Optimazation is done
This commit is contained in:
@@ -0,0 +1,349 @@
|
||||
/**++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
*
|
||||
* PANIC
|
||||
* Portable Algorithms and Numerics In C++
|
||||
*
|
||||
* Scientific computing from scratch, with feeling.
|
||||
*
|
||||
* Copyright (c) 2026 Michelle Bausager
|
||||
*
|
||||
* This file is part of PANIC.
|
||||
*
|
||||
* PANIC is free software licensed under the GNU General Public License v3.0 or later.
|
||||
* You may redistribute and/or modify it under the terms of the GPL.
|
||||
*
|
||||
* PANIC is distributed in the hope that it will be useful, but WITHOUT ANY WARRANTY;
|
||||
* without even the implied warranty of MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.
|
||||
* See the LICENSE file for the full license text.
|
||||
*
|
||||
* SPDX-License-Identifier: GPL-3.0-or-later
|
||||
*
|
||||
*++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
|
||||
*
|
||||
* Project Name: PANIC
|
||||
* Module Name: neural_network
|
||||
* File Name: optimizer_rmsprop.cpp
|
||||
* Revision: 0.1.0
|
||||
* Date: 04-08-2026
|
||||
* Author: Michelle Bausager
|
||||
*
|
||||
* Description:
|
||||
* Defines the optimizer_rmsprop used in neural network
|
||||
*
|
||||
*++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++*/
|
||||
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
// INCLUDE DESCRIPTION
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
#include <neural_network/optimizers/optimizer_rmsprop.hpp>
|
||||
#include <config/omp.hpp>
|
||||
|
||||
#include <math/mul.hpp>
|
||||
#include <math/add.hpp>
|
||||
#include <math/sub.hpp>
|
||||
#include <math/div.hpp>
|
||||
#include <math/sqrt.hpp>
|
||||
|
||||
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
// PRIVATE CONSTANTS
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
/**
|
||||
* @brief Minimum number of element operations before using the OpenMP-enabled loop.
|
||||
*
|
||||
* Small vectors and matrices are kept serial because the overhead of starting
|
||||
* worker threads can be larger than the work itself.
|
||||
*/
|
||||
static const panic::types::uint_t optimizer_rmsprop_omp_min_size = 250;
|
||||
|
||||
static const panic::types::real_t real_t_1 = static_cast<panic::types::real_t>(1);
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
// INPLEMENTATION
|
||||
//---------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
namespace panic{
|
||||
namespace neural_network{
|
||||
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Constructor Name : panic::neural_network::optimizer_rmsprop
|
||||
//
|
||||
// Description:
|
||||
// Constructor for optimizer_rmsprop.
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
optimizer_rmsprop::optimizer_rmsprop(const panic::types::real_t learning_rate,
|
||||
const panic::types::real_t decay,
|
||||
const panic::types::real_t epsilon,
|
||||
const panic::types::real_t rho){
|
||||
|
||||
this->learning_rate = learning_rate;
|
||||
this->decay = decay;
|
||||
this->epsilon = epsilon;
|
||||
this->rho = rho;
|
||||
|
||||
iterations = static_cast<panic::types::uint_t>(0);
|
||||
current_learning_rate = learning_rate;
|
||||
|
||||
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Constructor Name : panic::neural_network::update_params
|
||||
//
|
||||
// Description:
|
||||
// Updates weights and biases in layer.
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
bool optimizer_rmsprop::update_params(trainable_layer& layer) {
|
||||
|
||||
|
||||
// Gradients must match their corresponding parameters.
|
||||
if (layer.weights.rows() != layer.dweights.rows() || layer.weights.cols() != layer.dweights.cols() || layer.biases.size() != layer.dbiases.size()){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (layer.weights.rows() == static_cast<panic::types::uint_t>(0) || layer.weights.cols() == static_cast<panic::types::uint_t>(0) || layer.biases.size() == static_cast<panic::types::uint_t>(0) ){
|
||||
return false;
|
||||
}
|
||||
|
||||
|
||||
// If layer doesn't contain cache arrays, create them
|
||||
// and fill them with zeros
|
||||
if(layer.weights.rows() != layer.weight_cache.rows() ||
|
||||
layer.weights.cols() != layer.weight_cache.cols() ||
|
||||
layer.biases.size() != layer.bias_cache.size()){
|
||||
|
||||
if (! layer.weight_cache.resize(layer.weights.rows(), layer.weights.cols())){
|
||||
return false;
|
||||
}
|
||||
layer.weight_cache.fill(0);
|
||||
|
||||
if (! layer.bias_cache.resize(layer.biases.size())){
|
||||
return false;
|
||||
}
|
||||
layer.bias_cache.fill(0);
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Update caches with squared current gradients
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
const panic::types::real_t one_minus_rho =
|
||||
static_cast<panic::types::real_t>(1) - rho;
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Weight cache
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
panic::tensor::real_matrix retained_weight_cache;
|
||||
panic::tensor::real_matrix squared_dweights;
|
||||
panic::tensor::real_matrix weighted_dweights;
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.weight_cache,
|
||||
rho,
|
||||
retained_weight_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.dweights,
|
||||
layer.dweights,
|
||||
squared_dweights
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::mul(
|
||||
squared_dweights,
|
||||
one_minus_rho,
|
||||
weighted_dweights
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
retained_weight_cache,
|
||||
weighted_dweights,
|
||||
layer.weight_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Bias cache
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
panic::tensor::real_vector retained_bias_cache;
|
||||
panic::tensor::real_vector squared_dbiases;
|
||||
panic::tensor::real_vector weighted_dbiases;
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.bias_cache,
|
||||
rho,
|
||||
retained_bias_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.dbiases,
|
||||
layer.dbiases,
|
||||
squared_dbiases
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::mul(
|
||||
squared_dbiases,
|
||||
one_minus_rho,
|
||||
weighted_dbiases
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
retained_bias_cache,
|
||||
weighted_dbiases,
|
||||
layer.bias_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Calculate weight updates
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
panic::tensor::real_matrix weight_updates;
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.dweights,
|
||||
-current_learning_rate,
|
||||
weight_updates
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
panic::tensor::real_matrix sqrt_weight_cache;
|
||||
panic::tensor::real_matrix weight_denominator;
|
||||
|
||||
if (!panic::math::sqrt(
|
||||
layer.weight_cache,
|
||||
sqrt_weight_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
sqrt_weight_cache,
|
||||
epsilon,
|
||||
weight_denominator
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::div(
|
||||
weight_updates,
|
||||
weight_denominator,
|
||||
weight_updates
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
layer.weights,
|
||||
weight_updates,
|
||||
layer.weights
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Calculate bias updates
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
|
||||
panic::tensor::real_vector bias_updates;
|
||||
|
||||
if (!panic::math::mul(
|
||||
layer.dbiases,
|
||||
-current_learning_rate,
|
||||
bias_updates
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
panic::tensor::real_vector sqrt_bias_cache;
|
||||
panic::tensor::real_vector bias_denominator;
|
||||
|
||||
if (!panic::math::sqrt(
|
||||
layer.bias_cache,
|
||||
sqrt_bias_cache
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
sqrt_bias_cache,
|
||||
epsilon,
|
||||
bias_denominator
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::div(
|
||||
bias_updates,
|
||||
bias_denominator,
|
||||
bias_updates
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
if (!panic::math::add(
|
||||
layer.biases,
|
||||
bias_updates,
|
||||
layer.biases
|
||||
)){
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
|
||||
}
|
||||
|
||||
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Constructor Name : panic::neural_network::pre_update_params
|
||||
//
|
||||
// Description:
|
||||
// function to update internal parameters before update_params()
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
bool optimizer_rmsprop::pre_update_params(){
|
||||
|
||||
if (decay){
|
||||
current_learning_rate = learning_rate * (real_t_1 / (real_t_1 + (decay * iterations)));
|
||||
}
|
||||
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
// Constructor Name : panic::neural_network::post_update_params
|
||||
//
|
||||
// Description:
|
||||
// function to update internal parameters after update_params()
|
||||
//--------------------------------------------------------------------------------------------------------------------------
|
||||
bool optimizer_rmsprop::post_update_params(){
|
||||
|
||||
iterations += static_cast<panic::types::uint_t>(1);
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
} // namespace tensor
|
||||
} // namespace panic
|
||||
Reference in New Issue
Block a user