为何自定义DAG结构PyTorch神经网络无法正常训练?
我基于PyTorch实现了一个以神经元为单位而非层级运行的神经网络,定义了Neuron类和DAGNet类(代码如下),尝试训练DAGNet("A->D,E,F;B->D,E,F;C->D,E,F;D->G,H,I;E->G,H,I;F->G,H,I;G->J;H->J;I->J;J;"),该结构理论上等价于含3个3神经元隐藏层的标准神经网络,但损失无法降至0.33以下,而标准PyTorch网络训练效果良好,请问实现是否存在不等价之处?
定义的Neuron类
class Neuron(nn.Module): def __init__(self, n_inputs, name, activation=nn.ReLU()): self.name = name super(Neuron, self).__init__() self.linear = nn.Linear(n_inputs, 1) self.activation = activation nn.init.kaiming_uniform_(self.linear.weight, nonlinearity='relu') def forward(self, x): return self.activation(self.linear(x))
定义的DAGNet类
class DAGNet(nn.Module): def __init__(self, description_string, *args, **kwargs): """ Create a neural network based on a directed acyclic graph. @param description_string: A string describing the DAG. Example string: A->C,D,E;B->C,D,E;C->F,G,H;D->F,G,H;E->F,G,H;F->I,J,K;G->I,J,K;H->I,J,K;I->L;J->L;K->L;L; This implements a simple neural network with two inputs, three hidden layers, and one output layer. However, the DAG structure allows much more general networks to be constructed. """ super().__init__(*args, **kwargs) # First, create a standard directed acyclic graph with the information parsed from the string. # This will help us set up the actual neurons later. self.dag = {} for node in description_string.split(';'): if node == '': continue if '->' not in node: # Last node in the string has no outputs self.dag[node] = [] continue node_name, inputs = node.split('->') inputs = inputs.split(',') self.dag[node_name] = inputs # Reverse each edge in the topological sort to get the input edges for each neuron. self.input_edges = {} for node, edges in self.dag.items(): for edge in edges: if edge in self.input_edges: self.input_edges[edge].append(node) else: self.input_edges[edge] = [node] # Now, create a neuron for each node in the DAG. # Make sure to correctly count the number of inputs for each neuron. self.neurons = {} for node_name, inputs in self.dag.items(): n_inputs = 0 if node_name not in self.input_edges: n_inputs = 2 # (x, y) input neuron else: n_inputs += len(self.input_edges[node_name]) n_outputs = len(self.dag[node_name]) if len(self.dag[node_name]) > 0 else 1 self.neurons[node_name] = Neuron(n_inputs, n_outputs, node_name) # Topologically sort the neurons to ensure that the inputs to each neuron are computed before the neuron itself. self.ts = TopologicalSorter(self.dag) self.so = list(reversed(list(self.ts.static_order()))) self.model = nn.Sequential(*[self.neurons[node_name] for node_name in self.so]) def forward(self, x): # Run the input through the DAG to get the output. # Make sure to correctly pass the inputs to each neuron. node_names = list(self.neurons.keys()) outputs = {} # The input (x) gets loaded into nodes that don't have any inputs for node_name in self.so: if node_name not in self.input_edges: outputs[node_name] = self.neurons[node_name](x) # The rest of the nodes get their inputs from the outputs of other nodes for node_name in self.so: if node_name in self.input_edges: inputs = [] for input_node in self.input_edges[node_name]: # Get the subpart of the input node's output that corresponds to this node's input inputs.append(torch.hsplit(outputs[input_node], len(self.dag[input_node]))[self.dag[input_node].index(node_name)]) outputs[node_name] = self.neurons[node_name](torch.cat(inputs, dim=1)) return outputs[list(self.so)[-1]]
核心不等价点分析
1. Neuron类初始化参数传递完全错误
DAGNet中创建Neuron的代码:
self.neurons[node_name] = Neuron(n_inputs, n_outputs, node_name)
但Neuron的__init__签名是:
def __init__(self, n_inputs, name, activation=nn.ReLU()):
参数顺序完全混乱:
- 把
n_outputs传给了name参数,导致Neuron的名称变成了数字 - 把
node_name传给了activation参数,相当于用字符串作为激活函数,这会直接破坏forward逻辑 - 你定义的
n_outputs在Neuron类中完全没用——Neuron的线性层固定输出1,这个参数纯属多余
正确调用方式:
self.neurons[node_name] = Neuron(n_inputs, node_name)
2. 输入节点的输入维度硬编码,缺乏灵活性
在DAGNet初始化中,对无输入的节点(如A、B、C)直接设置:
if node_name not in self.input_edges: n_inputs = 2 # (x, y) input neuron
硬编码为2的输入维度,和标准PyTorch网络根据实际数据动态定义输入维度的逻辑不符,若输入数据维度不是2,会直接导致维度不匹配,影响收敛。
3. 神经元输出传递逻辑错误
forward中拆分输出的代码:
inputs.append(torch.hsplit(outputs[input_node], len(self.dag[input_node]))[self.dag[input_node].index(node_name)])
每个Neuron的输出是[batch_size, 1]的张量,用torch.hsplit按下游节点数量拆分是完全错误的——1列无法被拆分成多份,且单个神经元的输出应该直接传给所有下游节点,不需要拆分。
正确逻辑:
inputs.append(outputs[input_node])
直接取上游节点的完整输出,多个上游输出拼接后正好匹配下游Neuron的输入维度。
4. 拓扑排序顺序反转,导致计算逻辑混乱
代码中反转了拓扑排序结果:
self.so = list(reversed(list(self.ts.static_order())))
拓扑排序的正确顺序是前驱节点在当前节点之前,反转后会导致先计算下游节点、再计算上游节点的错误顺序,虽然你先处理了无输入节点,但整体计算逻辑已经混乱,影响网络的正常运行。
正确顺序:
self.so = list(self.ts.static_order())
5. 初始化方式与标准网络存在差异
你手动给Neuron的线性层设置了kaiming_uniform_初始化,而PyTorch默认的nn.Linear使用的是xavier_uniform初始化,虽然这不是核心问题,但会影响训练初期的收敛速度。
总结
上述错误中,参数传递错误和输出传递逻辑错误是最致命的,直接导致你的DAGNet结构和预期的标准网络完全不符,无法正常收敛。修正这些问题后,网络应该能达到和标准网络相近的训练效果。
内容的提问来源于stack exchange,提问作者416E64726577

