PyTorch Lightning调用trainer.test报'DummyTqdmFile'无'encoding'属性错误
PyTorch Lightning + Hyperopt: AttributeError on trainer.test()
问题详情
使用PyTorch Lightning训练模型并结合Hyperopt调参时,调用trainer.test()触发如下异常:
Traceback (most recent call last): File "/home/amaan/code/DL_Simulation/mtl_former.py", line 302, in <module> train() File "/home/amaan/code/DL_Simulation/mtl_former.py", line 111, in train best = fmin( File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/fmin.py", line 540, in fmin return trials.fmin( File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/base.py", line 671, in fmin return fmin( File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/fmin.py", line 586, in fmin rval.exhaust() File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/fmin.py", line 364, in exhaust self.run(self.max_evals - n_done, block_until_done=self.asynchronous) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/fmin.py", line 300, in run self.serial_evaluate() File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/fmin.py", line 178, in serial_evaluate result = self.domain.evaluate(spec, ctrl) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/hyperopt/base.py", line 892, in evaluate rval = self.fn(pyll_rval) File "/home/amaan/code/DL_Simulation/mtl_former.py", line 112, in <lambda> fn=lambda x: train_single(x, opt_run_id), File "/home/amaan/code/DL_Simulation/mtl_former.py", line 181, in train_single val_loss = model_transformer._training_model( File "/home/amaan/code/DL_Simulation/model_training/models/model_transformer.py", line 398, in _training_model final_val_loss = BaseTrainingMethods.training_suffix( File "/home/amaan/code/DL_Simulation/model_training/models/model_base.py", line 200, in training_suffix trainer.test( File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/trainer/trainer.py", line 753, in test return call._call_and_handle_interrupt( File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/trainer/call.py", line 44, in _call_and_handle_interrupt return trainer_fn(*args, **kwargs) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/trainer/trainer.py", line 793, in _test_impl results = self._run(model, ckpt_path=ckpt_path) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/trainer/trainer.py", line 986, in _run results = self._run_stage() File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/trainer/trainer.py", line 1023, in _run_stage return self._evaluation_loop.run() File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/loops/utilities.py", line 182, in _decorator return loop_run(self, *args, **kwargs) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/loops/evaluation_loop.py", line 142, in run return self.on_run_end() File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/loops/evaluation_loop.py", line 274, in on_run_end self._print_results(logged_outputs, self._stage.value) File "/home/amaan/.conda/envs/dl/lib/python3.9/site-packages/pytorch_lightning/loops/evaluation_loop.py", line 552, in _print_results if sys.stdout.encoding is not None: AttributeError: 'DummyTqdmFile' object has no attribute 'encoding'
注释掉trainer.test()后代码可正常运行,但在调试控制台单独执行该调用无报错。核心代码片段:
@staticmethod def training_suffix(model, args, output_folder, wandb_logger, datamodule, run): # 回调与Trainer初始化代码 early_stopping, checkpoint_callback, lr_monitor = ( CallbackStorage.get_default_callbacks( output_folder=output_folder, filename="model" ) ) args['epochs'] = 2 if args['debug'] else args['epochs'] trainer = pl.Trainer( fast_dev_run=False, max_epochs=args["epochs"], callbacks=[early_stopping, checkpoint_callback, lr_monitor], accelerator="gpu" if torch.cuda.is_available() else "cpu", logger=wandb_logger, check_val_every_n_epoch=1, ) trainer.fit(model, datamodule=datamodule) final_val_loss = float(trainer.callback_metrics.get("val_loss")) print("Final val_loss: ", final_val_loss) trainer.test( ckpt_path='best', dataloaders=datamodule.val_dataloader(), ) # wandb相关代码 wandb_reference = { 'entitiy' : run.entity, 'project' : run.project, 'id' : run.id } with open(os.path.join(output_folder, "wandb_information.json"), "w") as file: json.dump(wandb_reference, file) wandb.finish() return final_val_loss
问题原因
Hyperopt在运行时会用DummyTqdmFile重定向sys.stdout以适配进度条,但该对象没有PyTorch Lightning打印测试结果时需要的encoding属性,导致触发AttributeError。调试控制台中单独执行时,sys.stdout未被Hyperopt修改,因此无报错。
解决方案
方案1:临时补充缺失的encoding属性
在调用trainer.test()前,手动为当前sys.stdout添加encoding属性:
import sys # 在trainer.test()前插入 if not hasattr(sys.stdout, 'encoding'): sys.stdout.encoding = 'utf-8' trainer.test( ckpt_path='best', dataloaders=datamodule.val_dataloader(), )
方案2:禁用测试结果打印
通过参数关闭PyTorch Lightning的测试结果自动打印,避免触发属性检查:
# 调用trainer.test()时添加verbose=False trainer.test( ckpt_path='best', dataloaders=datamodule.val_dataloader(), verbose=False )
或者初始化Trainer时禁用相关打印:
trainer = pl.Trainer( # 原有参数保留 fast_dev_run=False, max_epochs=args["epochs"], callbacks=[early_stopping, checkpoint_callback, lr_monitor], accelerator="gpu" if torch.cuda.is_available() else "cpu", logger=wandb_logger, check_val_every_n_epoch=1, # 新增:禁用测试阶段的结果打印 enable_progress_bar=False, enable_test_loop_logging=False )
方案3:升级依赖库
检查并升级PyTorch Lightning和Hyperopt到最新稳定版本,官方可能已修复该兼容性问题:
pip install --upgrade pytorch-lightning hyperopt
内容的提问来源于stack exchange,提问作者Amaan Ansari
相关产品推荐
相关产品推荐

