加载TensorFlow训练模型时出现意外关键字参数'name'错误求助
加载TensorFlow训练模型时出现意外关键字参数'name'错误求助
我之前也遇到过完全一样的问题!这个错误的核心原因是:TensorFlow在加载自定义模型时,会自动把name参数传递给你的Agent类构造函数,但你的__init__方法并没有定义这个参数,所以直接报错了。下面给你两个亲测有效的解决办法:
方法一:修改Agent类的构造函数和配置方法(推荐)
只需要给__init__添加name参数并传递给父类,同时在get_config里把name加入配置项即可:
修改后的__init__方法:
def __init__(self, number_of_outputs: int, number_of_hidden_units: int, name=None): super(Agent,self).__init__(name=name) # 将name参数传递给父类构造函数 self.number_of_outputs = number_of_outputs self.number_of_hidden_units = number_of_hidden_units # 后续原有代码保持不变...
修改后的get_config方法:
def get_config(self): base_config = super().get_config() config = { "number_of_outputs": self.number_of_outputs, "number_of_hidden_units" :self.number_of_hidden_units, "name": self.name # 新增该行,将name纳入模型配置 } return {**base_config, **config}
修改完成后,建议重新保存一次模型再尝试加载,这样能确保模型配置的完整性,彻底解决加载时的参数问题。
方法二:加载时用包装类兼容(临时方案)
如果不想修改原Agent类的代码,可以在加载模型时用一个包装类来处理name参数:
def load_full_model(self, path_to_model): class WrappedAgent(Agent): def __init__(self, number_of_outputs, number_of_hidden_units, name=None): super().__init__(number_of_outputs, number_of_hidden_units) self.model = load_model(path_to_model, custom_objects={'Agent': WrappedAgent} )
不过这个方法只是临时绕过问题,还是推荐第一种方案,更符合TensorFlow的模型序列化规范。
补充说明:为什么会出现这个问题?
TensorFlow的Model基类本身就需要name参数来标识模型实例,当你保存模型时,这个name会被写入模型配置文件。加载时,框架会读取配置里的所有参数(包括name),并传递给自定义类的构造函数。你的原构造函数没声明接收这个参数,自然就会抛出“意外关键字参数”的错误了。
备注:内容来源于stack exchange,提问作者borjanob
相关产品推荐
相关产品推荐

