Stable-Baselines3加载模型出现未关闭文件ResourceWarning问题求助
问题描述
加载Stable Baselines3模型时,仅在使用自定义Gym环境时出现ResourceWarning: unclosed file <_io.BufferedReader name='saved_models\best_model.zip'>警告。模型训练、预测流程均正常,但重新加载模型时会抛出该警告。完整错误栈如下:
File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\matplotlibtools.py", line 30, in do_enable_gui enable_gui(guiname) File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\inputhook.py", line 540, in enable_gui return gui_hook(app) File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\inputhook.py", line 176, in enable_qt from pydev_ipython.qt_for_kernel import QT_API, QT_API_PYQT5 File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\qt_for_kernel.py", line 116, in <module> QtCore, QtGui, QtSvg, QT_API = load_qt(api_opts) File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\qt_loaders.py", line 276, in load_qt if not can_import(api): File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\qt_loaders.py", line 152, in can_import if not has_binding(api): File "c:\program files\microsoft visual studio\2022\community\common7\ide\extensions\microsoft\python\core\debugpy\_vendored\pydevd\pydev_ipython\qt_loaders.py", line 115, in has_binding import imp File "C:\Users\craig.evans\AppData\Local\Programs\Python\Python310\lib\imp.py", line 31, in <module> warnings.warn("the imp module is deprecated in favour of importlib and slated " DeprecationWarning: the imp module is deprecated in favour of importlib and slated for removal in Python 3.12; see the module's documentation for alternative uses Exception ignored in: <_io.FileIO name='saved_models\\best_model.zip' mode='rb' closefd=True> Traceback (most recent call last): File "C:\Users\craig.evans\AppData\Local\Programs\Python\Python310\lib\site-packages\stable_baselines3\common\base_class.py", line 659, in load data, params, pytorch_variables = load_from_zip_file( ResourceWarning: unclosed file <_io.BufferedReader name='saved_models\\best_model.zip'>
替换为官方环境(如LunarLanderContinuous-v2)时无此问题,可复现的代码如下:
import os from datetime import datetime from random import seed import gym import numpy as np import torch as th from stable_baselines3 import PPO from stable_baselines3.common.utils import set_random_seed from stable_baselines3.common.vec_env import SubprocVecEnv from stable_baselines3 import TD3 from stable_baselines3.common.monitor import Monitor from stable_baselines3.common.results_plotter import load_results, ts2xy from stable_baselines3.common.callbacks import BaseCallback platform = 1 # 1: surface book, 2: work machine num_env = 8 total_timesteps = 1_000_000 #10_000_000 # run this many total steps save_in_steps = 50_000 # save the network after this many steps number_training_steps = int(total_timesteps/save_in_steps) NSTEPS = 2000 VF_COEFF = 1.0 ENT_COEFF = 0.005 LEARNING_RATE = 0.0005 # changed from 0.0001 MINIBATCHES = 100 # changed from 128 EPOCHS = 5 results_folder = '.\\saved_models\\' run_name = 'Sim_' + datetime.now().strftime('%Y%m%d_%H%M%S') tensorboard_log_location = '.\\tensorboard\\' best_mean_reward, n_steps = -np.inf, 0 class SaveOnBestTrainingRewardCallback(BaseCallback): """ Callback for saving a model (the check is done every ``check_freq`` steps) based on the training reward (in practice, we recommend using ``EvalCallback``). :param check_freq: :param log_dir: Path to the folder where the model will be saved. It must contains the file created by the ``Monitor`` wrapper. :param verbose: Verbosity level: 0 for no output, 1 for info messages, 2 for debug messages """ def __init__(self, check_freq: int, log_dir: str, verbose: int = 1): super(SaveOnBestTrainingRewardCallback, self).__init__(verbose) self.check_freq = check_freq self.log_dir = log_dir self.save_path = os.path.join(log_dir, "best_model") self.best_mean_reward = -np.inf def _init_callback(self) -> None: # Create folder if needed if self.save_path is not None: os.makedirs(self.save_path, exist_ok=True) def _on_step(self) -> bool: if self.n_calls % self.check_freq == 0: # Retrieve training reward simResults = load_results(self.log_dir) try: x, y = ts2xy(simResults, "timesteps") if len(x) > 0: # Mean training reward over the last 100 episodes mean_reward = np.mean(y[-100:]) if self.verbose >= 1: print(f"Num timesteps: {self.num_timesteps}") print(f"Best mean reward: {self.best_mean_reward:.2f} - Last mean reward per episode: {mean_reward:.2f}") # New best model, you could save the agent here if mean_reward > self.best_mean_reward: self.best_mean_reward = mean_reward # Example for saving best model if self.verbose >= 1: print(f"Saving new best model to {self.save_path}") self.model.save(self.save_path) except: print('Error loading results') print(simResults) return True def make_env(env_id, rank: int, seed: int = 0, log_dir: str = ''): ''' Utility function for multiprocessed env. env_id: configuration information for the environment param rank: index of the subprocess param seed: the inital seed for RNG ''' def _init(): env = gym.make(env_id) env.seed(seed + rank) log_file = os.path.join(log_dir, str(rank)) if log_dir is not None else None return Monitor(env, log_file) set_random_seed(seed) return _init def get_policyNetwork()-> dict: '''https://stable-baselines3.readthedocs.io/en/sde/guide/custom_policy.html ''' policyNetwork = dict(activation_fn = th.nn.ReLU, net_arch = dict(vf=[256, 256], pi=[256, 256])) return policyNetwork if __name__ == '__main__': os.makedirs(results_folder, exist_ok=True) env = SubprocVecEnv([make_env( env_id = 'LunarLanderContinuous-v2', rank = i, seed = 0, log_dir = results_folder ) for i in range(num_env)]) customPolicy = get_policyNetwork() #https://medium.com/aureliantactics/ppo-hyperparameters-and-ranges-6fc2d29bccbe model = PPO( policy = 'MlpPolicy', env = env, verbose = 1, vf_coef = VF_COEFF, n_epochs = EPOCHS, ent_coef = ENT_COEFF, learning_rate = LEARNING_RATE, tensorboard_log = tensorboard_log_location, n_steps = NSTEPS, batch_size = MINIBATCHES, policy_kwargs = customPolicy, device = 'auto' ) saveCallback = SaveOnBestTrainingRewardCallback(check_freq = save_in_steps, log_dir = results_folder) savedModel = results_folder + 'best_model.zip' model.load(savedModel) model.learn(total_timesteps = total_timesteps, reset_num_timesteps = False, callback = saveCallback, progress_bar = True) inputResult = input('Would you like to render? (y/n): ') if inputResult =='y': env_render = gym.make('LunarLanderContinuous-v2') dones = False obs = env_render.reset() while dones == False: action, _states = model.predict(obs) obs, rewards, dones, info = env_render.step(action) env_render.render()
原因与解决方法
核心原因
该警告源于模型加载时,读取best_model.zip的文件流未被正确关闭。仅在自定义环境中出现,本质是自定义Gym环境的序列化/反序列化逻辑存在缺陷:Stable Baselines3保存模型时会将环境配置存入zip文件,加载时需反序列化这些内容。如果自定义环境中存在未正确处理的资源(如未关闭的文件句柄),或__getstate__/__setstate__方法实现不当,会导致zip文件流无法正常关闭。
具体解决方法
1. 修复自定义环境的序列化逻辑
确保自定义环境类实现正确的__getstate__和__setstate__方法,序列化时排除无法被pickle的资源:
class CustomEnv(gym.Env): def __init__(self): self.some_file = open("data.txt", "r") # 其他初始化逻辑 def __getstate__(self): # 序列化时移除未关闭的文件句柄 state = self.__dict__.copy() del state['some_file'] return state def __setstate__(self, state): self.__dict__.update(state) # 反序列化时重新初始化必要资源 self.some_file = open("data.txt", "r")
2. 手动管理模型加载的文件流
不直接使用model.load(),而是手动打开并关闭zip文件流,避免资源泄漏:
from stable_baselines3.common.save_util import load_from_zip_file savedModel = results_folder + 'best_model.zip' with open(savedModel, "rb") as f: data, params, pytorch_variables = load_from_zip_file(f) model = PPO.load_from_vectorized_params(params=params, data=data, pytorch_variables=pytorch_variables, env=env)
3. 升级Stable Baselines3版本
部分旧版本的Stable Baselines3在处理自定义环境序列化时存在bug,升级到最新稳定版可解决:
pip install --upgrade stable-baselines3
4. 临时禁用资源警告(不推荐)
若不影响功能,仅需消除警告,可在代码开头添加:
import warnings warnings.filterwarnings("ignore", category=ResourceWarning)
此方法仅治标,建议优先解决序列化逻辑问题。
内容的提问来源于stack exchange,提问作者Craig Evans
相关产品推荐
相关产品推荐

