在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
相关产品推荐
相关产品推荐

