-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathConstantOptimizer.cs
More file actions
63 lines (50 loc) · 2.26 KB
/
Copy pathConstantOptimizer.cs
File metadata and controls
63 lines (50 loc) · 2.26 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
using System.Collections.Generic;
using UnityEngine;
public class ConstantOptimizer
{
private float learningRate = 0.01f;
private int maxIterations = 100;
private float tolerance = 1e-6f;
public void OptimizeConstants(Individual individual, float[] inputData, float[] outputData)
{
List<ExpressionNode> constants = new List<ExpressionNode>();
CollectConstants(individual.root, constants);
if (constants.Count == 0) return;
for (int iter = 0; iter < maxIterations; iter++)
{
float currentMSE = CalculateMSE(individual.root, inputData, outputData);
foreach (ExpressionNode constant in constants)
{
float originalValue = constant.constantValue;
constant.constantValue = originalValue + tolerance;
float msePlus = CalculateMSE(individual.root, inputData, outputData);
constant.constantValue = originalValue - tolerance;
float mseMinus = CalculateMSE(individual.root, inputData, outputData);
float gradient = (msePlus - mseMinus) / (2 * tolerance);
constant.constantValue = originalValue - learningRate * gradient;
}
float newMSE = CalculateMSE(individual.root, inputData, outputData);
if (Mathf.Abs(newMSE - currentMSE) < tolerance)
break;
}
individual.CalculateFitness(inputData, outputData, individual.complexity);
}
private float CalculateMSE(ExpressionNode root, float[] inputData, float[] outputData)
{
float mse = 0f;
for (int i = 0; i < inputData.Length; i++)
{
float predicted = root.Evaluate(inputData[i]);
float error = outputData[i] - predicted;
mse += error * error;
}
return mse / inputData.Length;
}
private void CollectConstants(ExpressionNode node, List<ExpressionNode> constants)
{
if (node == null) return;
if (node.nodeType == NodeType.Constant) constants.Add(node);
CollectConstants(node.left, constants);
CollectConstants(node.right, constants);
}
}