# How to get model summary using Pytorch backend?

**URL:** <https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064>\
**Category:** RLlib\
**Created:** [May 6, 2021, 8:05am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064 "2021-05-06T08:05:31Z")\
**Posts on this page:** 7\
**Page:** 1

<div class="post-metadata">

**Author:** ![bug404](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/bug404/32/931_2.png) [@bug404](https://discuss.ray.io/u/bug404)\
**Post date:** [May 6, 2021, 8:05am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/1 "2021-05-06T08:05:31Z")

</div>

If use tf2, the model.summary() of keras can help output the model summary, and the rllib also defines the base\_model. But if use Pytorch, how to output the model summary?

---

<div class="post-metadata">

**Author:** ![sven1977](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/sven1977/32/53_2.png) [@sven1977](https://discuss.ray.io/u/sven1977)\
**Post date:** [May 6, 2021, 8:46am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/2 "2021-05-06T08:46:03Z")

</div>

Hey @bug404 . If you have an RLlib Trainer object:

```auto
print(trainer.get_policy().model)

```

If you are using tune.run, try to add this line into ray/rllib/policy/torch\_policy.py:~164 (after we assign self.model to the 0th of the multi-GPU towers):

```auto
print(self.model)

```

---

<div class="post-metadata">

**Author:** ![bug404](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/bug404/32/931_2.png) [@bug404](https://discuss.ray.io/u/bug404)\
**Post date:** [May 6, 2021, 9:03am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/3 "2021-05-06T09:03:43Z")

</div>

It’s very helpful, thank you very much.

---

<div class="post-metadata">

**Author:** ![bug404](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/bug404/32/931_2.png) [@bug404](https://discuss.ray.io/u/bug404)\
**Post date:** [May 6, 2021, 9:29am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/4 "2021-05-06T09:29:03Z")

</div>

It’s a very cool way modifying the source code to support this feature, haha.

---

<div class="post-metadata">

**Author:** ![Glaucus-2G](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/glaucus-2g/32/1103_2.png) [@Glaucus-2G](https://discuss.ray.io/u/Glaucus-2G)\
**Post date:** [June 17, 2021, 12:16pm UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/5 "2021-06-17T12:16:07Z")

</div>

Hey @sven1977 ,  
I use `policy.model.base_model.summary()` to output shape of model, but it reports  
`AttributeError: 'FullyConnectedNetwork' object has no attribute 'base_model'` .

So I just use `trainer.get_policy().model` to see it, and output is:

```auto
FullyConnectedNetwork(
  (_logits): SlimFC(
    (_model): Sequential(
      (0): Linear(in_features=32, out_features=6, bias=True)
    )
  )
  (_hidden_layers): Sequential(
    (0): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=20, out_features=32, bias=True)
        (1): ReLU()
      )
    )
    (1): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=32, out_features=64, bias=True)
        (1): ReLU()
      )
    )
    (2): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=64, out_features=32, bias=True)
        (1): ReLU()
      )
    )
  )
  (_value_branch_separate): Sequential(
    (0): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=20, out_features=32, bias=True)
        (1): ReLU()
      )
    )
    (1): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=32, out_features=64, bias=True)
        (1): ReLU()
      )
    )
    (2): SlimFC(
      (_model): Sequential(
        (0): Linear(in_features=64, out_features=32, bias=True)
        (1): ReLU()
      )
    )
  )
  (_value_branch): SlimFC(
    (_model): Sequential(
      (0): Linear(in_features=32, out_features=1, bias=True)
    )
  )
)

```

The input latitude of my environment is 20 and the output latitude is 3. And the definition of environment and network initialization are as follows:

```auto
class Env(gym.Env)
    def __init__ (self):
        self.action_space = Box(-1, 1, [3,])
        self.observation_space = Box(-1, 1, [20,])
    ...

------------------
    ray.init()
    config = DEFAULT_CONFIG.copy()
    config['model']['fcnet_hiddens'] = [32, 64, 32]
    config['model']['fcnet_activation'] = "relu"
    ...

```

So I’m a little confused about that the `out_features` of network is not 3.

---

<div class="post-metadata">

**Author:** ![mannyv](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/mannyv/32/606_2.png) [@mannyv](https://discuss.ray.io/u/mannyv)\
**Post date:** [June 17, 2021, 2:20pm UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/6 "2021-06-17T14:20:12Z")

</div>

Hi @Glaucus-2G,

If you take a look at this [dataflow diagram](https://docs.ray.io/en/master/_images/rllib-components.svg), the part you are describing is the green model box on the right part of the image. The outputs of the model are the logits that are sent to the ActionDistribution box. It is the ActionDistribution box that converts those logits into the actual actions.

In your example you will get 32 logit outputs from the model and 3 action outputs from the ActionDistribution. Does this make sense?

---

<div class="post-metadata">

**Author:** ![Glaucus-2G](https://sea2.discourse-cdn.com/flex020/user_avatar/discuss.ray.io/glaucus-2g/32/1103_2.png) [@Glaucus-2G](https://discuss.ray.io/u/Glaucus-2G)\
**Post date:** [June 23, 2021, 6:31am UTC](https://discuss.ray.io/t/how-to-get-model-summary-using-pytorch-backend/2064/7 "2021-06-23T06:31:48Z")

</div>

Hey @mannyv ,  
Thank you for your reply！You helped me understand this content very well. I have ignored the function of the ActionDistribution box before.
