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

如何无需存储训练数据加载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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 20:20:55