如何无需存储训练数据加载KrigingSurrogate已训练模型?
跳过OpenMDAO KrigingSurrogate的训练数据校验直接加载模型
OpenMDAO的KrigingSurrogate支持通过training_cache参数缓存训练好的代理模型,但加载时会强制校验用户提供的训练数据与缓存中的数据是否一致,不一致就会重新训练。这导致跨脚本复用模型时,必须额外序列化存储训练数据,十分繁琐。
有没有办法跳过这个校验步骤,直接使用缓存模型中已保存的训练数据?
当前跨脚本创建与加载模型的写法
创建模型脚本(create_model.py)
import numpy as np import openmdao.api as om import pickle x = np.arange(0, 11, 1) y = x**2 surrogate = om.MetaModelUnStructuredComp() surrogate.add_input('x', training_data=x) surrogate.add_output('y', training_data=y, surrogate=om.KrigingSurrogate(training_cache='surrogate.dat')) prob = om.Problem() prob.model.add_subsystem('surrogate', surrogate) prob.setup() prob.run_model() # 训练模型并保存到surrogate.dat # 额外序列化训练数据供加载脚本使用 training_data = {'x': x, 'y': y} with open('training_data.pickle', 'wb') as f: pickle.dump(training_data, f)
加载模型脚本(load_model.py)
import numpy as np import openmdao.api as om import pickle # 必须加载序列化的训练数据用于校验,否则无法加载缓存模型 with open('training_data.pickle', 'rb') as f: training_data = pickle.load(f) x = training_data['x'] y = training_data['y'] surrogate = om.MetaModelUnStructuredComp() surrogate.add_input('x', training_data=x) surrogate.add_output('y', training_data=y, surrogate=om.KrigingSurrogate(training_cache='surrogate.dat')) prob = om.Problem() prob.model.add_subsystem('surrogate', surrogate) prob.setup() prob.run_model() # 加载缓存模型
解决方案:自定义代理类跳过校验
通过继承KrigingSurrogate自定义子类,重写训练逻辑,跳过数据校验步骤,直接加载缓存中的模型和训练数据。
自定义代理类代码
import pickle import openmdao.api as om from openmdao.surrogate_models.kriging import KrigingSurrogate class NoCheckKrigingSurrogate(KrigingSurrogate): def _train(self): # 优先尝试加载缓存 if self.training_cache is not None: try: with open(self.training_cache, 'rb') as f: self.X, self.Y, self.krg = pickle.load(f) # 直接使用缓存内的训练数据,跳过校验 return except (FileNotFoundError, pickle.UnpicklingError): # 加载失败时走正常训练流程 pass # 原类的正常训练逻辑 super()._train() # 训练完成后保存缓存 if self.training_cache is not None: with open(self.training_cache, 'wb') as f: pickle.dump((self.X, self.Y, self.krg), f)
修改后的加载脚本(load_model.py)
import numpy as np import openmdao.api as om import pickle from your_module import NoCheckKrigingSurrogate # 替换为自定义类所在模块 # 从缓存读取训练数据维度,用于定义输入输出的形状 with open('surrogate.dat', 'rb') as f: X, Y, _ = pickle.load(f) surrogate = om.MetaModelUnStructuredComp() surrogate.add_input('x', shape=X.shape[1]) # 匹配输入维度 surrogate.add_output('y', shape=Y.shape[1], surrogate=NoCheckKrigingSurrogate(training_cache='surrogate.dat')) prob = om.Problem() prob.model.add_subsystem('surrogate', surrogate) prob.setup() prob.run_model() # 直接加载缓存模型 # 测试预测 prob.set_val('surrogate.x', 5) prob.run_model() print(prob.get_val('surrogate.y')) # 输出:[25.]
内容的提问来源于stack exchange,提问作者awass
相关产品推荐
相关产品推荐

