# Explain this code - Pytorch CNN

**URL:** <https://discuss.pytorch.org/t/explain-this-code-pytorch-cnn/120544>\
**Category:** Uncategorized\
**Created:** [May 7, 2021, 8:59am UTC](https://discuss.pytorch.org/t/explain-this-code-pytorch-cnn/120544 "2021-05-07T08:59:50Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![tomi\_datasci](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/tomi_datasci/32/30122_2.png) [@tomi\_datasci](https://discuss.pytorch.org/u/tomi_datasci)\
**Post date:** [May 7, 2021, 8:59am UTC](https://discuss.pytorch.org/t/explain-this-code-pytorch-cnn/120544/1 "2021-05-07T08:59:51Z")

</div>

These two lines of code are from the testing of a CNN model. I know of alternative ways of getting the predictions and correct predictions  
but I have struggled to make sense of the two lines below:

```
  pred = output.data.max(1, keepdim=True)[1]
  correct += pred.eq(target.data.view_as(pred)).sum()
```

---

<div class="post-metadata">

**Author:** ![Sanjan\_Das](https://discuss.pytorch.org/user_avatar/discuss.pytorch.org/sanjan_das/32/36341_2.png) [@Sanjan\_Das](https://discuss.pytorch.org/u/Sanjan_Das)\
**Post date:** [May 8, 2021, 2:44pm UTC](https://discuss.pytorch.org/t/explain-this-code-pytorch-cnn/120544/2 "2021-05-08T14:44:34Z")

</div>

So at a high level - what your code is doing is first getting the predicted tensor `pred`, and then element-wise comparing them to the values in tensor `target`, setting them to `True` if the elements match and `False` if not. And then when you take the `sum()`, you’re simply summing over the `True` values, and that gives you the number of correct predictions.
