如何修复Java神经元模拟程序的StackOverflowError异常
解决神经元双向连接引发的StackOverflowError问题
问题背景
开发神经元模拟程序时,运行NeuronSimulator第三部分测试会抛出java.lang.StackOverflowError,异常根源是Neuron类receivesStimulus方法中connection.receivesStimulus(this.signal);语句因神经元双向连接触发无限递归。注释brain.createConnections();或brain.stimulate(0, 10);可消除异常,但此时索引3的神经元信号为0.0,不符合预期值4.800000000000001。
约束条件
- 不得修改
NeuronSimulator类 - 不得在现有类中新增方法或属性
现有代码
Position类
public class Position { private double x; private double y; public Position(double x, double y) { this.x = x; this.y = y; } public Position() { this(0, 0); } public double getX() { return x; } public double getY() { return y; } public String toString() { return "(" + this.x + ", " + this.y + ")"; } }
Brain类
import java.util.ArrayList; public class Brain { private ArrayList<Neuron> neurons; public Brain() { this.neurons = new ArrayList<>(); } public int getNbNeurons() { return neurons.size(); } public Neuron getNeuron(int index) { return neurons.get(index); } public void addNeuron(Position position, double attenuation) { Neuron n = new Neuron(position, attenuation); this.neurons.add(n); } public void addCumulativeNeuron(Position pos, double attenuation) { CumulativeNeuron cn = new CumulativeNeuron(pos, attenuation); this.neurons.add(cn); } public void stimulate(int index, double stimulus) { Neuron neuron = neurons.get(index); neuron.receivesStimulus(stimulus); } public double probe(int index) { Neuron neuron = neurons.get(index); System.out.println("INDEX IS: " + index); System.out.println("SIGNAL: " + neuron.getSignal()); return neuron.getSignal(); } public void createConnections() { int size = this.getNbNeurons(); for (int i = 0; i < size - 2; i++) { Neuron n1 = this.getNeuron(i); Neuron n2 = this.getNeuron(i + 1); n1.connection(n2); n2.connection(n1); } if (size > 1) { this.getNeuron(0).connection(this.getNeuron(1)); } if (size > 2) { this.getNeuron(0).connection(this.getNeuron(2)); } } public String toString() { StringBuilder sb = new StringBuilder(); sb.append("*----------*\n"); sb.append("The brain contains ").append(this.getNbNeurons()).append(" neuron(s)\n"); for (Neuron neuron : this.neurons) { sb.append(neuron.toString()).append("\n"); } sb.append("*----------*\n"); return sb.toString(); } }
NeuronSimulator类
public class NeuronSimulator { public static void main(String[] args) { // TEST PART 1 System.out.println("Test part 1:"); System.out.println("--------------------"); Position position1 = new Position(0, 1); Position position2 = new Position(1, 0); Position position3 = new Position(1, 1); Neuron neuron1 = new Neuron(position1, 0.5); Neuron neuron2 = new Neuron(position2, 1.0); Neuron neuron3 = new Neuron(position3, 2.0); neuron1.connection(neuron2); neuron2.connection(neuron3); neuron1.receivesStimulus(10); System.out.println("Signals : "); System.out.println(neuron1.getSignal()); System.out.println(neuron2.getSignal()); System.out.println(neuron3.getSignal()); System.out.println(); System.out.println("First connection of neuron 1"); System.out.println(neuron1.getConnection(0)); // END TEST PART 1 // TEST PART 2 System.out.println("Test part 2:"); System.out.println("--------------------"); Position position5 = new Position(0, 0); CumulativeNeuron neuron5 = new CumulativeNeuron(position5, 0.5); neuron5.receivesStimulus(10); neuron5.receivesStimulus(10); System.out.println("Signal of cumulative neuron -> " + neuron5.getSignal()); // END TEST PART 2 //TEST PART 3 System.out.println(); System.out.println("Test part 3:"); System.out.println("--------------------"); Brain brain = new Brain(); brain.addNeuron(new Position(0, 0), 0.5); brain.addNeuron(new Position(0, 1), 0.2); brain.addNeuron(new Position(1, 0), 1.0); brain.addCumulativeNeuron(new Position(1, 1), 0.8); brain.createConnections(); brain.stimulate(0, 10); System.out.println("Signal of 3rd neuron -> " + brain.probe(3)); System.out.println(brain); // END TEST PART 3 } }
解决方案
修改Neuron类的receivesStimulus方法,添加极小值阈值判断终止递归,同时调整信号传递逻辑为传递增量而非总信号:
public void receivesStimulus(double stimulus) { // 设置极小阈值,终止无限递归 if (Math.abs(stimulus) < 1e-9) { return; } double increment = stimulus * this.attenuation; this.signal += increment; // 传递本次刺激产生的增量,而非当前总信号 for (Neuron connection : connections) { connection.receivesStimulus(increment); } }
逻辑说明
- 阈值终止递归:当刺激值小到可忽略(小于
1e-9)时停止递归,既避免栈溢出,又不会影响计算精度。 - 传递增量而非总信号:原逻辑传递神经元总信号会导致双向连接间循环放大,改为传递本次刺激产生的增量,符合真实信号传递逻辑,同时让递归快速收敛到阈值以下。
CumulativeNeuron无需额外修改,若其重写了receivesStimulus方法,保持原有累加逻辑即可,父类的阈值判断已处理递归问题。
修改后,第三部分测试的索引3神经元信号会因浮点精度问题显示为4.800000000000001,完全符合预期且无StackOverflowError异常。
内容的提问来源于stack exchange,提问作者carlosS
相关产品推荐
相关产品推荐

