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

Skorch封装GPyTorch ExactGP:多次.fit()调用参数未重置求助

问题描述

使用Skorch封装GPyTorch的ExactGP模型,通过子类化ExactGPRegressor适配业务场景,模型训练正常,但在Jupyter Notebook不同单元格调用.fit()时,模型未重新初始化,仍从之前状态继续训练——即便已设置warm_start=False并限制训练轮数。

尝试在fit()方法中调用.initialize()但无效,以下是自定义封装类的简化版本:

class ExactGPModel(gpytorch.models.ExactGP):
    def __init__(self, train_inputs, train_targets, likelihood, covar_module):
        super().__init__(train_inputs, train_targets, likelihood=likelihood)
        self.mean_module = gpytorch.means.ConstantMean()
        self.covar_module = covar_module

    def forward(self, X):
        mean_x = self.mean_module(X)
        covar_x = self.covar_module(X)
        return gpytorch.distributions.MultivariateNormal(mean_x, covar_x)


class CustomGPWrapper(ExactGPRegressor):
    def __init__(self, module, likelihood, criterion, dtype=torch.float32, **kwargs):
        self.module = module
        self.likelihood = likelihood 
        self.criterion = criterion
        self.dtype = dtype
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        super().__init__(module=module, **kwargs)

    def fit(self, X, y, **fit_params):
        X_array = np.array(X) if not isinstance(X, np.ndarray) else X
        y_array = np.array(y).flatten()
        X_tensor = torch.tensor(X_array.astype(self.dtype))
        y_tensor = torch.tensor(y_array)

        # 尝试强制重新初始化
        self.initialize()

        return super().fit(X_tensor, y_tensor, **fit_params)

尽管如此设置,模型参数仍在多次.fit()调用间保留,训练延续上次状态而非从头开始。需要解决:如何确保每次.fit()调用时模型完全重新初始化?有没有更优封装方式避免状态保留?

解决方案

1. 核心问题分析

ExactGP模型在初始化时会绑定训练数据(train_inputs、train_targets),且这些数据会作为模型状态的一部分被保留。Skorch的initialize()方法仅重置模型参数,不会重新绑定新的训练数据,也不会清除ExactGP内部的训练状态缓存,导致参数无法彻底重置。

2. 修复方案:每次fit重新创建ExactGP实例

修改CustomGPWrapper,在每次调用.fit()时用当前训练数据重新实例化ExactGPModel,同时重置Skorch的内部状态:

class CustomGPWrapper(ExactGPRegressor):
    def __init__(self, model_cls, likelihood, criterion, covar_module, dtype=torch.float32, **kwargs):
        # 传入模型类和组件,而非预创建的实例
        self.model_cls = model_cls
        self.likelihood = likelihood 
        self.criterion = criterion
        self.covar_module = covar_module
        self.dtype = dtype
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        # 初始化时传入临时占位module,后续fit中替换
        super().__init__(module=lambda: None, **kwargs)

    def fit(self, X, y, **fit_params):
        X_array = np.array(X) if not isinstance(X, np.ndarray) else X
        y_array = np.array(y).flatten()
        X_tensor = torch.tensor(X_array.astype(self.dtype)).to(self.device)
        y_tensor = torch.tensor(y_array).to(self.device)

        # 用当前训练数据重新创建ExactGP实例
        self.module_ = self.model_cls(
            train_inputs=X_tensor,
            train_targets=y_tensor,
            likelihood=self.likelihood,
            covar_module=self.covar_module
        ).to(self.device)
        
        # 更新criterion的绑定模型
        if isinstance(self.criterion, gpytorch.mlls.ExactMarginalLogLikelihood):
            self.criterion.model = self.module_
        
        # 重置Skorch的优化器、调度器等状态
        self.initialize()
        
        return super().fit(X_tensor, y_tensor, **fit_params)

3. 关键修改说明

  • 初始化时不再传入预构建的ExactGPModel实例,改为传入模型类model_cls和协方差模块covar_module,保留组件复用性
  • 每次fit时,用当前训练数据重新创建ExactGPModel并赋值给self.module_(Skorch内部存储模型的属性)
  • 更新ExactMarginalLogLikelihood的绑定模型,确保损失计算正确
  • 调用self.initialize()重置Skorch的优化器、学习率调度器等状态,保证从头开始训练

4. 使用示例

# 定义固定组件
likelihood = gpytorch.likelihoods.GaussianLikelihood()
covar_module = gpytorch.kernels.ScaleKernel(gpytorch.kernels.RBFKernel())
criterion = gpytorch.mlls.ExactMarginalLogLikelihood(likelihood, None)

# 创建Wrapper实例
gp_wrapper = CustomGPWrapper(
    model_cls=ExactGPModel,
    likelihood=likelihood,
    criterion=criterion,
    covar_module=covar_module,
    warm_start=False,
    max_epochs=100,
    optimizer=torch.optim.Adam,
    optimizer__lr=0.1
)

# 第一次训练
gp_wrapper.fit(X_train1, y_train1)

# 第二次训练(完全从头初始化)
gp_wrapper.fit(X_train2, y_train2)

5. 额外优化建议

  • 若需重复使用同一组超参数,可将模型组件(如likelihood、covar_module)的初始化逻辑封装成函数,减少重复代码
  • 在Jupyter Notebook中,若多次运行同一单元格创建Wrapper实例,确保每次都重新实例化CustomGPWrapper,避免复用旧变量的残留状态

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.12 19:45:16