使用TensorBoard可视化PyTorch网络遇NumPy2.0错误,求临时修复
问题:TensorBoard与NumPy 2.0兼容问题的临时解决办法
我想用TensorBoard将PyTorch构建的神经网络以图的形式可视化,代码如下:
PyTorch网络代码
import torch BATCH_SIZE = 16 DIM_IN = 1000 HIDDEN_SIZE = 100 DIM_OUT = 10 class TinyModel(torch.nn.Module): def __init__(self): super(TinyModel, self).__init__() self.layer1 = torch.nn.Linear(DIM_IN, HIDDEN_SIZE) self.relu = torch.nn.ReLU() self.layer2 = torch.nn.Linear(HIDDEN_SIZE, DIM_OUT) def forward(self, x): x = self.layer1(x) x = self.relu(x) x = self.layer2(x) return x some_input = torch.randn(BATCH_SIZE, DIM_IN, requires_grad=False) ideal_output = torch.randn(BATCH_SIZE, DIM_OUT, requires_grad=False) model = TinyModel()
TensorBoard配置代码
from torch.utils.tensorboard import SummaryWriter # Create a SummaryWriter writer = SummaryWriter("checkpoint") # Add the graph to TensorBoard writer.add_graph(model, some_input) writer.close()
运行tensorboard --logdir=checkpoint时出现以下错误:
Traceback (most recent call last): File "/home/k/python_venv/bin/tensorboard", line 5, in <module> from tensorboard.main import run_main File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/main.py", line 27, in <module> from tensorboard import default File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/default.py", line 39, in <module> from tensorboard.plugins.hparams import hparams_plugin File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/plugins/hparams/hparams_plugin.py", line 30, in <module> from tensorboard.plugins.hparams import backend_context File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/plugins/hparams/backend_context.py", line 26, in <module> from tensorboard.plugins.hparams import metadata File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/plugins/hparams/metadata.py", line 32, in <module> NULL_TENSOR = tensor_util.make_tensor_proto( File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/util/tensor_util.py", line 405, in make_tensor_proto numpy_dtype = dtypes.as_dtype(nparray.dtype) File "/home/k/python_venv/lib/python3.10/site-packages/tensorboard/compat/tensorflow_stub/dtypes.py", line 677, in as_dtype if type_value.type == np.string_ or type_value.type == np.unicode_: File "/home/k/python_venv/lib/python3.10/site-packages/numpy/__init__.py", line 397, in __getattr__ raise AttributeError( AttributeError: `np.string_` was removed in the NumPy 2.0 release. Use `np.bytes_` instead.. Did you mean: 'strings'?
该问题可能在后续版本中修复,但目前有以下可行的临时解决办法:
降级NumPy到1.x版本
直接回退到NumPy 1.x系列,避开2.0的API变更,执行命令:pip install numpy<2.0这是最稳妥的方案,等TensorBoard官方发布兼容NumPy 2.0的版本后再升级。
手动修改TensorBoard的兼容代码
找到报错的文件/home/k/python_venv/lib/python3.10/site-packages/tensorboard/compat/tensorflow_stub/dtypes.py,定位到第677行左右的代码:if type_value.type == np.string_ or type_value.type == np.unicode_:将其修改为:
if type_value.type == np.bytes_ or type_value.type == np.str_:注意:后续TensorBoard更新会覆盖这个修改,需要在更新后重新操作。
安装TensorBoard开发版本
如果官方已经在主分支修复了该问题,可以直接安装最新的开发版本:pip install --upgrade git+https://github.com/tensorflow/tensorboard.git此方案可能存在其他不稳定因素,适合愿意尝试新版本的用户。
内容的提问来源于stack exchange,提问作者pkj
相关产品推荐
相关产品推荐

