Generic 1-layer training

This commit is contained in:
Pascal Serrarens 2026-07-03 11:44:58 +02:00
parent c16c3c5dc9
commit a38e4c3c77
2 changed files with 50 additions and 48 deletions

View File

@ -706,65 +706,65 @@ namespace NanoBrain {
Debug.Log($"Updated weight: {loss.magnitude} {loss} {scaledOutput} {synapse.weight}"); Debug.Log($"Updated weight: {loss.magnitude} {loss} {scaledOutput} {synapse.weight}");
} }
public void BackPropagation1(Vector3 cost, Vector3 error, float learningRate) { // public void BackPropagation1(Vector3 cost, Vector3 error, float learningRate) {
cost = Vector3.Scale(error, error); // error^2 // cost = Vector3.Scale(error, error); // error^2
float3 derivative = 2 * error; // derivative of (error^2) // float3 derivative = 2 * error; // derivative of (error^2)
// inverted because it uses the non-convential // // inverted because it uses the non-convential
// error=(actual-taget) instead of (target-actual) // // error=(actual-taget) instead of (target-actual)
// dSSR / dPredicted // // dSSR / dPredicted
// Bias // // Bias
float3 deltaBias = derivative; // float3 deltaBias = derivative;
// deltaBias *= 1; // because bias is always fully applied // // deltaBias *= 1; // because bias is always fully applied
Vector3 stepSize = deltaBias * learningRate; // Vector3 stepSize = deltaBias * learningRate;
this.bias -= stepSize; // this.bias -= stepSize;
foreach (Synapse synapse in this.synapses) { // foreach (Synapse synapse in this.synapses) {
// derivative for the weight? // // derivative for the weight?
float3 deltaSynapse = derivative; // dSSR/dPredicted // float3 deltaSynapse = derivative; // dSSR/dPredicted
// derivative for the previous activation // // derivative for the previous activation
deltaSynapse *= synapse.neuron.activation; // dPredicted/dWeight // deltaSynapse *= synapse.neuron.activation; // dPredicted/dWeight
// // derivative for the activator // // // derivative for the activator
// switch (activator) { // // switch (activator) {
// case ActivationType.Linear: // // case ActivationType.Linear:
// //delta2 *= 1; // // //delta2 *= 1;
// break; // // break;
// default: // // default:
// break; // // break;
// } // // }
float deltaWeight = length(deltaSynapse); // float deltaWeight = length(deltaSynapse);
synapse.weight += learningRate * deltaWeight; // synapse.weight += learningRate * deltaWeight;
synapse.neuron.BackPropagation2(derivative * synapse.weight, learningRate); // synapse.neuron.BackPropagation2(derivative * synapse.weight, learningRate);
} // }
} // }
public void BackPropagation0(Vector3 error, float learningRate) { public void BackPropagation0(float error, float learningRate) {
float3 derivative = 2 * error; // derivative of (error^2) float derivative = 2 * error; // derivative of (error^2)
// inverted because it uses the non-convential // inverted because it uses the non-convential
// error=(actual-taget) instead of (target-actual) // error=(actual-taget) instead of (target-actual)
// dSSR / dPredicted // dSSR / dPredicted
BackPropagation2(derivative, learningRate); BackPropagation2(derivative, learningRate);
} }
public void BackPropagation2(Vector3 derivative, float learningRate) { public void BackPropagation2(float derivative, float learningRate) {
// Bias // Bias
float3 deltaBias = derivative; // dSSR/dActivator // float3 deltaBias = derivative; // dSSR/dActivator
switch (activator) { // dActivator/dBias // switch (activator) { // dActivator/dBias
case ActivationType.Linear: // case ActivationType.Linear:
//deltaBias *= 1; // //deltaBias *= 1;
break; // break;
default: // default:
break; // break;
} // }
// deltaBias *= 1; // because bias is always fully applied // // deltaBias *= 1; // because bias is always fully applied
Vector3 stepSize = deltaBias * learningRate; // Vector3 stepSize = deltaBias * learningRate;
this.bias -= stepSize; // this.bias -= stepSize;
foreach (Synapse synapse in this.synapses) { foreach (Synapse synapse in this.synapses) {
synapse.BackPropagation(length(derivative), learningRate); synapse.BackPropagation(this, derivative, learningRate);
// // derivative for the weight? // // derivative for the weight?
// float3 deltaSynapse = derivative; // dSSR/dActivator // float3 deltaSynapse = derivative; // dSSR/dActivator
@ -785,7 +785,7 @@ namespace NanoBrain {
// float deltaWeight = length(deltaSynapse); // float deltaWeight = length(deltaSynapse);
// synapse.weight += learningRate * deltaWeight; // synapse.weight += learningRate * deltaWeight;
BackPropagation2(derivative * synapse.weight, learningRate); // BackPropagation2(derivative * synapse.weight, learningRate);
} }
} }

View File

@ -33,13 +33,15 @@ namespace NanoBrain {
this.weight = weight; this.weight = weight;
} }
public virtual void BackPropagation(float error, float learningRate) { public virtual void BackPropagation(Neuron receiver, float error, float learningRate) {
float derivative = error; float derivative = error;
switch (neuron.activator) {
switch (receiver.activator) {
case Neuron.ActivationType.Linear: case Neuron.ActivationType.Linear:
derivative *= 1; derivative *= 1;
break; break;
default: default:
Debug.Log("other activator");
break; break;
} }
derivative *= math.length(neuron.activation); derivative *= math.length(neuron.activation);