diff --git a/Runtime/Scripts/Core/Neuron.cs b/Runtime/Scripts/Core/Neuron.cs index cb5639c..b5ebda2 100644 --- a/Runtime/Scripts/Core/Neuron.cs +++ b/Runtime/Scripts/Core/Neuron.cs @@ -6,8 +6,7 @@ using Unity.Mathematics; using static Unity.Mathematics.math; #endif -namespace NanoBrain -{ +namespace NanoBrain { /// /// A neuron is a basic Nucleus @@ -22,20 +21,17 @@ namespace NanoBrain /// Each connection has a weight which is used to multiply the output of that other neuron /// before it is used by the combinator. [Serializable] - public class Neuron : Nucleus - { + public class Neuron : Nucleus { /// /// Create a new Neuron in a Cluster instance /// /// The parent cluster in which the new Neuron should be created /// The name of the new Neuron - public Neuron(Cluster parent, string name) - { + public Neuron(Cluster parent, string name) { this.parent = parent; this.name = name; - if (this.parent != null) - { + if (this.parent != null) { this.parent.nuclei ??= new(); this.parent.nuclei.Add(this); } @@ -67,8 +63,7 @@ namespace NanoBrain /// The weight applied to the input. Default value = 1 /// The created Synapse /// This will add a new input to this nucleus with the given weight. - public Synapse AddSynapse(Neuron sendingNucleus, float weight = 1) - { + public Synapse AddSynapse(Neuron sendingNucleus, float weight = 1) { Synapse synapse = new(sendingNucleus, weight); this.synapses.Add(synapse); return synapse; @@ -79,8 +74,7 @@ namespace NanoBrain /// /// The sender of the input to the Synapse /// The found Synapse or null when the sender has no synapse to this nucleus. - public Synapse GetSynapse(Nucleus sender) - { + public Synapse GetSynapse(Nucleus sender) { foreach (Synapse synapse in this.synapses) if (synapse.neuron == sender) return synapse; @@ -91,8 +85,7 @@ namespace NanoBrain /// Remove a synapse from a Nucleus /// /// Remote the synapse connecting to this Nucleus - public void RemoveSynapse(Nucleus sendingNucleus) - { + public void RemoveSynapse(Nucleus sendingNucleus) { this.synapses.RemoveAll(synapse => synapse.neuron == sendingNucleus); } @@ -102,8 +95,7 @@ namespace NanoBrain /// Set the bias, recalculate the output and update all Nuclei receiving from this Nucleus /// /// - public virtual void SetBias(Vector3 inputValue) - { + public virtual void SetBias(Vector3 inputValue) { this.bias = inputValue; this.lastUpdate = Time.time; this.parent?.UpdateFromNucleus(this); @@ -113,8 +105,7 @@ namespace NanoBrain /// The type of combinators /// /// A combinator combines the weighted values of the synapses to a single value - public enum CombinatorType - { + public enum CombinatorType { /// Add the weighted values together Sum, /// Multiply the weighted values @@ -129,8 +120,7 @@ namespace NanoBrain /// /// The type of /// - public enum ActivationType - { + public enum ActivationType { Linear, Power, Sqrt, @@ -149,11 +139,9 @@ namespace NanoBrain /// /// The activation funtion /// - public ActivationType activator - { + public ActivationType activator { get { return _activator; } - set - { + set { _activator = value; //this.curve = GenerateCurve(); } @@ -171,16 +159,15 @@ namespace NanoBrain /// /// The output value of the neuron /// - public virtual float3 outputValue - { + public virtual float3 outputValue { get { return _outputValue; } - set - { + set { _outputValue = value; if (this.isFiring) WhenFiring?.Invoke(); } } + public float3 activation => outputValue; /// /// The magnitude of the neuron output /// @@ -238,10 +225,8 @@ namespace NanoBrain /// Check if the neuron is sleeping. /// /// This will reset the output value if it is sleeping - public void SleepCheck() - { - if (this.isSleeping && this.outputSqrMagnitude > 0) - { + public void SleepCheck() { + if (this.isSleeping && this.outputSqrMagnitude > 0) { #if UNITY_MATHEMATICS this._outputValue = new float3(0, 0, 0); #else @@ -267,10 +252,8 @@ namespace NanoBrain public bool breakOnUpdate = false; /// \copydoc NanoBrain::Nucleus::ShallowCloneTo - public override Nucleus ShallowCloneTo(Cluster parent) - { - Neuron clone = new(parent, this.name) - { + public override Nucleus ShallowCloneTo(Cluster parent) { + Neuron clone = new(parent, this.name) { // prefabNucleus = this }; CloneFields(clone); @@ -281,8 +264,7 @@ namespace NanoBrain /// Copy relevant fields of this neuron to the given neuron /// /// - protected virtual void CloneFields(Neuron clone) - { + protected virtual void CloneFields(Neuron clone) { clone.bias = this.bias; clone.persistOutput = this.persistOutput; clone.combinator = this.combinator; @@ -294,45 +276,34 @@ namespace NanoBrain /// Delete the give neuron /// /// The neuron to delete - public static void Delete(Nucleus nucleus) - { + public static void Delete(Nucleus nucleus) { if (nucleus == null) return; - if (nucleus is Neuron neuron) - { - foreach (Synapse synapse in neuron.synapses) - { - if (synapse.neuron is Neuron synapse_nucleus) - { - if (synapse_nucleus.receivers.Count > 1) - { + if (nucleus is Neuron neuron) { + foreach (Synapse synapse in neuron.synapses) { + if (synapse.neuron is Neuron synapse_nucleus) { + if (synapse_nucleus.receivers.Count > 1) { // there is another nucleus feeding into this input nucleus synapse_nucleus.receivers.RemoveAll(r => r == nucleus); } - else - { + else { // No other links, delete it. Neuron.Delete(synapse_nucleus); } } } - foreach (Nucleus receiver in neuron.receivers) - { + foreach (Nucleus receiver in neuron.receivers) { if (receiver is not Neuron receiverNeuron) continue; if (receiver != null && receiverNeuron.synapses != null) receiverNeuron.synapses.RemoveAll(s => s.neuron == nucleus); } } - else if (nucleus is Cluster cluster) - { + else if (nucleus is Cluster cluster) { // remove all receivers for this cluster - foreach (Nucleus clusterNucleus in cluster.nuclei) - { - if (clusterNucleus is Neuron output) - { - foreach (Nucleus receiver in output.receivers) - { + foreach (Nucleus clusterNucleus in cluster.nuclei) { + if (clusterNucleus is Neuron output) { + foreach (Nucleus receiver in output.receivers) { if (receiver is not Neuron receiverNeuron) continue; receiverNeuron.synapses.RemoveAll(s => s.neuron == output); @@ -341,22 +312,19 @@ namespace NanoBrain } } - if (nucleus.parent.prefab != null) - { + if (nucleus.parent.prefab != null) { nucleus.parent.nuclei.RemoveAll(n => n == nucleus); nucleus.parent.RefreshOutputs(); } } /// \copydoc NanoBrain::Nucleus::UpdateStateIsolated - public override void UpdateStateIsolated() - { - if (breakOnUpdate) - { + public override void UpdateStateIsolated() { + if (breakOnUpdate) { Debug.Break(); } - this.combinationValue = Combinator(this.bias, this.synapses); - this.outputValue = Activator(this.combinationValue); + this.combination = Combinator(this.bias, this.synapses); + this.outputValue = Activator(this.combination); this.lastUpdate = Time.time; } @@ -364,7 +332,7 @@ namespace NanoBrain #if UNITY_MATHEMATICS [NonSerialized] - public float3 combinationValue; + public float3 combination; /// /// The combinator which combines the bias with the values from all synapses @@ -372,10 +340,8 @@ namespace NanoBrain /// The bias of the neuron /// The synapses of the neuron /// - protected float3 Combinator(float3 bias, List synapses) - { - switch (combinator) - { + protected float3 Combinator(float3 bias, List synapses) { + switch (combinator) { case CombinatorType.Sum: return CombinatorSum(bias, synapses); case CombinatorType.Product: @@ -391,11 +357,9 @@ namespace NanoBrain /// The bias of the neuron /// The synapses of the neuron /// - public static float3 CombinatorSum(float3 bias, List synapses) - { + public static float3 CombinatorSum(float3 bias, List synapses) { float3 sum = bias; - foreach (Synapse synapse in synapses) - { + foreach (Synapse synapse in synapses) { synapse.neuron.SleepCheck(); sum += synapse.weight * synapse.neuron.outputValue; } @@ -408,11 +372,9 @@ namespace NanoBrain /// The bias of the neuron /// The synapses of the neuron /// The result of the multiplication - public static float3 CombinatorProduct(float3 bias, List synapses) - { + public static float3 CombinatorProduct(float3 bias, List synapses) { float3 product = bias; - foreach (Synapse synapse in synapses) - { + foreach (Synapse synapse in synapses) { synapse.neuron.SleepCheck(); product *= synapse.weight * synapse.neuron.outputValue; } @@ -479,10 +441,8 @@ namespace NanoBrain /// /// The result of applying the activation function // This does not allocate memory and seems faster than a switch expression - protected float3 Activator(float3 inputValue) - { - switch (activator) - { + protected float3 Activator(float3 inputValue) { + switch (activator) { case ActivationType.Linear: return ActivatorLinear(inputValue); case ActivationType.Sqrt: @@ -507,8 +467,7 @@ namespace NanoBrain /// /// Input value /// The unchanged value - protected float3 ActivatorLinear(float3 input) - { + protected float3 ActivatorLinear(float3 input) { return input; } @@ -517,8 +476,7 @@ namespace NanoBrain /// /// Input value /// The square root of the input - protected float3 ActivatorSqrt(float3 input) - { + protected float3 ActivatorSqrt(float3 input) { float3 result = normalize(input) * MathF.Sqrt(length(input)); return result; } @@ -528,8 +486,7 @@ namespace NanoBrain /// /// Input value /// The input to the power of 2 - protected float3 ActivatorPower(float3 input) - { + protected float3 ActivatorPower(float3 input) { float3 result = normalize(input) * MathF.Pow(length(input), 2); return result; } @@ -539,8 +496,7 @@ namespace NanoBrain /// /// Input value /// 1/input value - protected float3 ActivatorReciprocal(float3 input) - { + protected float3 ActivatorReciprocal(float3 input) { float magnitude = length(input); if (magnitude == 0) return new float3(0, 0, 0); @@ -554,8 +510,7 @@ namespace NanoBrain /// /// Input value /// Tanh(input value) - protected float3 ActivatorTanh(float3 input) - { + protected float3 ActivatorTanh(float3 input) { float magnitude = length(input); float3 result = normalize(input) * MathF.Tanh(magnitude); return result; @@ -565,8 +520,7 @@ namespace NanoBrain /// /// Input value /// An uniform vector with magnitude between 0 and 1 - protected float3 ActivatorBinary(float3 input) - { + protected float3 ActivatorBinary(float3 input) { float magnitude = length(input); float value = Mathf.Clamp01(magnitude); return float3(value, value, value); @@ -577,8 +531,7 @@ namespace NanoBrain /// /// Input value /// The normalized vector - protected float3 ActivatorNormalized(float3 input) - { + protected float3 ActivatorNormalized(float3 input) { if (lengthsq(input) == 0) return input; float3 result = normalize(input); @@ -664,8 +617,7 @@ namespace NanoBrain /// /// The nuclei which have a synapse to this neuron /// - public virtual List receivers - { + public virtual List receivers { get { return _receivers; } set { _receivers = value; } } @@ -675,8 +627,7 @@ namespace NanoBrain /// /// The receiver to add /// The weight to use for the synapse to his neuron - public virtual void AddReceiver(Nucleus receiverToAdd, float weight = 1) - { + public virtual void AddReceiver(Nucleus receiverToAdd, float weight = 1) { if (receiverToAdd is not Neuron receiverNeuron) return; this._receivers.Add(receiverNeuron); @@ -689,8 +640,7 @@ namespace NanoBrain /// Remove a receiver to this neuron /// /// The receiver to remove - public virtual void RemoveReceiver(Nucleus receiverToRemove) - { + public virtual void RemoveReceiver(Nucleus receiverToRemove) { if (receiverToRemove is not Neuron receiverNeuron) return; this._receivers.RemoveAll(receiver => receiver == receiverNeuron); @@ -708,8 +658,7 @@ namespace NanoBrain #region Back propagation - public void BackPropagation(Synapse synapse, Vector3 error, float learningRate) - { + public void BackPropagation(Synapse synapse, Vector3 error, float learningRate) { // Loss function: // Mean Squared Error (MSE) 1/n * sum(errors^2) // We use simplified here 1/2 * (error^2) @@ -720,17 +669,16 @@ namespace NanoBrain // Backpropagation = loss * d(combinator) Vector3 delta2; - switch (activator) - { + switch (activator) { case ActivationType.Linear: // Derivative of this (f'()) would be 1. delta2 = loss * 1; break; case ActivationType.Power: - delta2 = loss * (2 * this.combinationValue); + delta2 = loss * (2 * this.combination); break; case ActivationType.Reciprocal: - delta2 = loss * (-1 / (this.combinationValue * this.combinationValue)); + delta2 = loss * (-1 / (this.combination * this.combination)); break; default: delta2 = loss; @@ -743,43 +691,108 @@ namespace NanoBrain Debug.Log($"Updated weight: {error.magnitude} {error} {scaledOutput} {synapse.weight}"); } - public void BackPropagationWithLoss(Synapse synapse, Vector3 loss, float learningRate) - { - // loss is a derivative of error - // Backpropagation = loss * d(combinator) - - Vector3 delta2; - switch (activator) - { - case ActivationType.Linear: - // Derivative of this (f'()) would be 1. - delta2 = loss * 1; - break; - case ActivationType.Power: - delta2 = loss * (2 * this.combinationValue); - break; - case ActivationType.Reciprocal: - delta2 = loss * (-1 / (this.combinationValue * this.combinationValue)); - break; - default: - delta2 = loss; - break; - } - + public void BackPropagationWithLoss(Synapse synapse, Vector3 loss, float learningRate) { + Vector3 delta2 = activator switch { + ActivationType.Linear => loss * 1, + ActivationType.Power => (Vector3)(loss * (2 * this.combination)), + ActivationType.Reciprocal => (Vector3)(loss * (-1 / (this.combination * this.combination))), + _ => loss, + }; Vector3 scaledOutput = Vector3.Scale(delta2, synapse.neuron.outputValue); float deltaWeight = Mathf.Abs(scaledOutput.x) + Mathf.Abs(scaledOutput.y) + Mathf.Abs(scaledOutput.z); synapse.weight += learningRate * deltaWeight; Debug.Log($"Updated weight: {loss.magnitude} {loss} {scaledOutput} {synapse.weight}"); } + public void BackPropagation1(Vector3 cost, Vector3 error, float learningRate) { + cost = Vector3.Scale(error, error); // error^2 + float3 derivative = 2 * error; // derivative of (error^2) + // inverted because it uses the non-convential + // error=(actual-taget) instead of (target-actual) + // dSSR / dPredicted + + // Bias + float3 deltaBias = derivative; + // deltaBias *= 1; // because bias is always fully applied + Vector3 stepSize = deltaBias * learningRate; + this.bias -= stepSize; + + foreach (Synapse synapse in this.synapses) { + // derivative for the weight? + float3 deltaSynapse = derivative; // dSSR/dPredicted + // derivative for the previous activation + deltaSynapse *= synapse.neuron.activation; // dPredicted/dWeight + + // // derivative for the activator + // switch (activator) { + // case ActivationType.Linear: + // //delta2 *= 1; + // break; + // default: + // break; + // } + float deltaWeight = length(deltaSynapse); + synapse.weight += learningRate * deltaWeight; + + synapse.neuron.BackPropagation2(derivative * synapse.weight, learningRate); + } + } + + public void BackPropagation0(Vector3 error, float learningRate) { + float3 derivative = 2 * error; // derivative of (error^2) + // inverted because it uses the non-convential + // error=(actual-taget) instead of (target-actual) + // dSSR / dPredicted + BackPropagation2(derivative, learningRate); + } + + public void BackPropagation2(Vector3 derivative, float learningRate) { + + // Bias + float3 deltaBias = derivative; // dSSR/dActivator + switch (activator) { // dActivator/dBias + case ActivationType.Linear: + //deltaBias *= 1; + break; + default: + break; + } + // deltaBias *= 1; // because bias is always fully applied + Vector3 stepSize = deltaBias * learningRate; + this.bias -= stepSize; + + foreach (Synapse synapse in this.synapses) { + // derivative for the weight? + float3 deltaSynapse = derivative; // dSSR/dActivator + + // derivative for the activator + // dActivator/dCombinator + switch (activator) { + case ActivationType.Linear: + //deltaSynapse *= 1; + break; + default: + break; + } + + // derivative for the previous activation + // dCombinator/dWeight + deltaSynapse *= synapse.neuron.activation; + + float deltaWeight = length(deltaSynapse); + synapse.weight += learningRate * deltaWeight; + + BackPropagation2(derivative * synapse.weight, learningRate); + } + } + #endregion Back propagation /// /// Process an external stimulus /// /// The value of the stimulus - public virtual void ProcessStimulus(Vector3 inputValue) - { + public virtual void ProcessStimulus(Vector3 inputValue) { this.lastUpdate = Time.time; this.bias = inputValue; this.parent?.UpdateFromNucleus(this);