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

在Keras中能否通过名称直接获取模型中的对应权重?

Keras按名称获取层内权重的实现方案

我们已知可以通过model.get_layer("layer_name")方法获取模型model的指定层对象,目前Keras官方未内置直接按名称提取层内权重的方法,但可以通过简单的自定义封装实现你期望的调用效果。


实现逻辑

层对象的weights属性存储了该层所有已定义的权重变量,每个变量都带有name属性。由于Keras会自动为权重名拼接层名作为前缀(例如lstm_1层下定义的recurrent_kernel最终名称为lstm_1/recurrent_kernel:0),我们只需提取变量名的最后一段匹配自定义的权重名即可。


具体实现方案

方案1:独立工具函数

def get_variable_by_name(layer, target_var_name):
    for weight in layer.weights:
        # 切分变量名取最后一段,匹配定义时使用的权重名
        pure_var_name = weight.name.rsplit("/", 1)[-1].split(":")[0]
        if pure_var_name == target_var_name:
            return weight
    raise ValueError(f"层{layer.name}中未找到名为{target_var_name}的权重")

调用方式:

recurrent_kernel = get_variable_by_name(model.get_layer("layer_name"), "recurrent_kernel")

方案2:扩展Layer类实现链式调用

如果需要完全匹配你期望的model.get_layer("layer_name").get_variable_by_name("recurrent_kernel")调用格式,可以给Keras的Layer基类动态添加自定义方法:

import tensorflow as tf

def get_variable_by_name(self, target_var_name):
    for weight in self.weights:
        pure_var_name = weight.name.rsplit("/", 1)[-1].split(":")[0]
        if pure_var_name == target_var_name:
            return weight
    raise ValueError(f"层{self.name}中未找到名为{target_var_name}的权重")

tf.keras.layers.Layer.get_variable_by_name = get_variable_by_name

扩展完成后即可直接使用你预期的写法调用:

recurrent_kernel = model.get_layer("layer_name").get_variable_by_name("recurrent_kernel")

内容的提问来源于stack exchange,提问作者Homero Esmeraldo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 19:24:02