无法使用PyTorch打印DQN模型摘要?求解决办法
解决DQN模型使用torchinfo打印模型摘要的报错问题
问题背景
为CartPole任务构建DQN强化学习模型时,想要实现类似Keras中model.summary()的功能,使用torchinfo.summary()时出现报错。
原始DQN模型代码
class DQN(): ''' Deep Q Neural Network class. ''' def __init__(self, state_dim, action_dim, hidden_dim=64, lr=0.05): super(DQN, self).__init__() self.criterion = torch.nn.MSELoss() self.model = torch.nn.Sequential( torch.nn.Linear(state_dim, hidden_dim), torch.nn.ReLU(), torch.nn.Linear(hidden_dim, hidden_dim*2), torch.nn.ReLU(), torch.nn.Linear(hidden_dim*2, action_dim) ) self.optimizer = torch.optim.Adam(self.model.parameters(), lr) def update(self, state, y): """Update the weights of the network given a training sample. """ y_pred = self.model(torch.Tensor(state)) loss = self.criterion(y_pred, Variable(torch.Tensor(y))) self.optimizer.zero_grad() loss.backward() self.optimizer.step() def predict(self, state): """ Compute Q values for all actions using the DQL. """ with torch.no_grad(): return self.model(torch.Tensor(state))
模型实例化代码
# Number of states = 4 n_state = env.observation_space.shape[0] # Number of actions = 2 n_action = env.action_space.n # Number of episodes episodes = 150 # Number of hidden nodes in the DQN n_hidden = 50 # Learning rate lr = 0.001 simple_dqn = DQN(n_state, n_action, n_hidden, lr)
尝试的torchinfo调用代码
from torchinfo import summary simple_dqn = DQN(n_state, n_action, n_hidden, lr) summary(simple_dqn, input_size=(4, 2, 50))
报错信息
NotImplementedError Traceback (most recent call last) /usr/local/lib/python3.7/dist-packages/torchinfo/torchinfo.py in forward_pass(model, x, batch_dim, cache_forward_pass, device, mode, **kwargs) 286 if isinstance(x, (list, tuple)): --> 287 _ = model.to(device)(*x, **kwargs) 288 elif isinstance(x, dict): 4 frames /usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _call_impl(self, *input, **kwargs) 1147 --> 1148 result = forward_call(*input, **kwargs) 1149 if _global_forward_hooks or self._forward_hooks: /usr/local/lib/python3.7/dist-packages/torch/nn/modules/module.py in _forward_unimplemented(self, *input) 200 """ --> 201 raise NotImplementedError(f"Module [{type(self).__name__}] is missing the required \"forward\" function") 202 NotImplementedError: Module [DQN] is missing the required "forward" function The above exception was the direct cause of the following exception: RuntimeError Traceback (most recent call last) <ipython-input-24-ee921f7e5cb5> in <module> 1 from torchinfo import summary 2 simple_dqn = DQN(n_state, n_action, n_hidden, lr) ----> 3 summary(simple_dqn, input_size=(4, 2, 50)) /usr/local/lib/python3.7/dist-packages/torchinfo/torchinfo.py in summary(model, input_size, input_data, batch_dim, cache_forward_pass, col_names, col_width, depth, device, dtypes, mode, row_settings, verbose, **kwargs) 216 ) 217 summary_list = forward_pass( --> 218 model, x, batch_dim, cache_forward_pass, device, model_mode, **kwargs 219 ) 220 formatting = FormattingOptions(depth, verbose, columns, col_width, rows) /usr/local/lib/python3.7/dist-packages/torchinfo/torchinfo.py in forward_pass(model, x, batch_dim, cache_forward_pass, device, mode, **kwargs) 297 "Failed to run torchinfo. See above stack traces for more details. " 298 f"Executed layers up to: {executed_layers}" --> 299 ) from e 300 finally: 301 if hooks: RuntimeError: Failed to run torchinfo. See above stack traces for more details. Executed layers up to: []
解决方案
报错核心原因有两点:
- 自定义
DQN类未继承torch.nn.Module,也未实现PyTorch模块必需的forward方法,导致torchinfo无法执行前向传播统计模型信息。 input_size参数设置错误,需匹配模型输入的维度(CartPole状态维度为4,需加上批量维度)。
修改后的DQN代码
import torch from torch import nn from torch.autograd import Variable class DQN(nn.Module): # 继承nn.Module ''' Deep Q Neural Network class. ''' def __init__(self, state_dim, action_dim, hidden_dim=64, lr=0.05): super(DQN, self).__init__() self.criterion = nn.MSELoss() self.model = nn.Sequential( nn.Linear(state_dim, hidden_dim), nn.ReLU(), nn.Linear(hidden_dim, hidden_dim*2), nn.ReLU(), nn.Linear(hidden_dim*2, action_dim) ) self.optimizer = torch.optim.Adam(self.model.parameters(), lr) def forward(self, x): # 添加forward方法 return self.model(x) def update(self, state, y): """Update the weights of the network given a training sample. """ y_pred = self.forward(torch.Tensor(state)) # 改用forward方法 loss = self.criterion(y_pred, Variable(torch.Tensor(y))) self.optimizer.zero_grad() loss.backward() self.optimizer.step() def predict(self, state): """ Compute Q values for all actions using the DQL. """ with torch.no_grad(): return self.forward(torch.Tensor(state)) # 改用forward方法
正确调用torchinfo的代码
from torchinfo import summary # 实例化模型 simple_dqn = DQN(n_state, n_action, n_hidden, lr) # input_size设置为(批量大小, 状态维度),这里批量大小设为1,状态维度是4 summary(simple_dqn, input_size=(1, 4))
执行后即可输出类似Keras的模型摘要,包含各层输出形状、参数数量等关键信息。
内容的提问来源于stack exchange,提问作者Collander
相关产品推荐
相关产品推荐

