Keras自定义层嵌套子层权重未纳入trainable_weights如何解决
问题原因
子层参数无法被trainable_weights收集的核心原因是Keras的自动参数追踪机制没有正常触发,具体有两个错误:
- 你直接给
self.layers赋值为普通Python列表,覆盖了Keras Layer类自带的内置layers属性——这个属性本来是框架用来自动存储所有被追踪的子层的,被你覆盖之后,框架的自动追踪逻辑直接失效。 - 就算换变量名,原生Python列表里存储的子层默认不会被Keras自动递归追踪,框架不会主动遍历普通列表去收集里面的子层和子层的嵌套参数。
另外你当前代码还有两个隐藏问题:
- GTConv层里
self.bias = None,实际并没有创建对应的可训练偏置参数,你之前计算的17个参数总数是不准的,bias部分实际不存在。 - GTN类没有实现
call方法,前向传播逻辑缺失,就算参数被追踪到,计算图里没有对应参数的计算路径,梯度也会返回None。
修复步骤
1. 修正子层存储逻辑,触发Keras自动追踪
把自定义子层列表的变量名换掉,不要占用内置的self.layers名称,同时让Keras可以追踪到列表里的所有子层,修改GTN的__init__方法:
class GTN(layers.Layer): def __init__(self, num_edge, num_channels, w_in, w_out, num_class,num_layers,norm): super(GTN, self).__init__() self.num_channels = num_channels self.w_in = w_in self.w_out = w_out self.num_class = num_class self.num_layers = num_layers # 换自定义变量名,不要覆盖内置self.layers self.gt_layers = [] # 初始化子层,这里用原生Python range即可,不需要tf.range for i in range(num_layers): if i == 0: layer = GTLayer(num_edge, num_channels, first=True) else: layer = GTLayer(num_edge, num_channels, first=False) self.gt_layers.append(layer) # 关键:调用Keras接口追踪列表内的所有子层 self.gt_layers = tf.keras.utils.track_list(self.gt_layers) w_init = tf.random_normal_initializer() self.weight = tf.Variable(initial_value= w_init(shape=(w_in, w_out)),trainable=True)
如果你不想调用track_list,也可以在append子层的时候,逐一把子层设为GTN的直接属性,比如setattr(self, f'gt_layer_{i}', layer),一样可以触发自动追踪。
2. 补全GTN的call方法,打通前向传播路径
在GTN类里实现call方法,确保前向传播时所有子层都被正确调用,参数能进入计算图,参考框架如下,你可以替换成自己的实际业务逻辑:
def call(self, A, node_features, train_node, train_target, training=False): H = None Ws = [] # 逐层执行GTLayer前向计算 for gt_layer in self.gt_layers: if gt_layer.first: H, W = gt_layer(A) else: H, W = gt_layer(A, H) Ws.extend(W) # 补全你自己的后续计算逻辑:特征变换、分类、损失计算 # 示例:用GTN自身的weight做特征投影 projected_feat = tf.matmul(node_features, self.weight) # ... 后续计算得到loss、预测输出y_train return loss, y_train, Ws
另外注意GTLayer的call方法现在只处理了first=True的分支,first=False的分支逻辑也要补全,不然执行到非首层会报错。
3. 修复GTConv的偏置定义(如果需要bias参数)
如果你确实需要GTConv层的bias参数,不要把self.bias设为None,和weight一样创建可训练变量即可:
class GTConv(keras.layers.Layer): def __init__(self, in_channels, out_channels): super(GTConv, self).__init__() w_init = tf.random_normal_initializer() self.weight = tf.Variable( initial_value=w_init(shape=(out_channels,in_channels,1,1)), trainable=True) # 按需创建bias参数 self.bias = tf.Variable( initial_value=w_init(shape=(out_channels,)), trainable=True ) self.scale = tf.Variable([0.1] , trainable=False) def call(self, A): A = tf.reduce_sum(A*(tf.nn.softmax(self.weight,1)), 1) # 按需加偏置计算 return A
验证方法
模型初始化完成后,直接打印len(model.trainable_weights)核对参数数量:
- 不含bias的情况:1个GTN自身weight + num_layers2个GTConv的weight,比如你代码里num_layers=3的话,总共有1+32=7个可训练参数。
- 含bias的情况:每个GTConv多1个bias参数,总参数为1 + num_layers22,num_layers=3时就是13个,和打印结果一致就说明参数追踪正常。
内容的提问来源于stack exchange,提问作者Anshuman Sinha
相关产品推荐
相关产品推荐

