You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

无法使用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: []

解决方案

报错核心原因有两点:

  1. 自定义DQN类未继承torch.nn.Module,也未实现PyTorch模块必需的forward方法,导致torchinfo无法执行前向传播统计模型信息。
  2. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.12 07:50:56