You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Keras自定义层嵌套子层权重未纳入trainable_weights如何解决

问题原因

子层参数无法被trainable_weights收集的核心原因是Keras的自动参数追踪机制没有正常触发,具体有两个错误:

  • 你直接给self.layers赋值为普通Python列表,覆盖了Keras Layer类自带的内置layers属性——这个属性本来是框架用来自动存储所有被追踪的子层的,被你覆盖之后,框架的自动追踪逻辑直接失效。
  • 就算换变量名,原生Python列表里存储的子层默认不会被Keras自动递归追踪,框架不会主动遍历普通列表去收集里面的子层和子层的嵌套参数。

另外你当前代码还有两个隐藏问题:

  1. GTConv层里self.bias = None,实际并没有创建对应的可训练偏置参数,你之前计算的17个参数总数是不准的,bias部分实际不存在。
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.26 13:54:22