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

