# Flux.jl confusion matrix

**URL:** <https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740>\
**Category:** General Usage\
**Tags:** flux\
**Created:** [January 17, 2019, 9:48am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740 "2019-01-17T09:48:56Z")\
**Posts on this page:** 14\
**Page:** 1

<div class="post-metadata">

**Author:** ![essenciary](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/essenciary/32/210469_2.png) [@essenciary](https://discourse.julialang.org/u/essenciary)\
**Post date:** [January 17, 2019, 9:48am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/1 "2019-01-17T09:48:56Z")

</div>

I’m using Flux to implement a handwritten digit classifier based on the MNIST dataset (ex [https://github.com/FluxML/model-zoo/blob/master/vision/mnist/mlp.jl](https://github.com/FluxML/model-zoo/blob/master/vision/mnist/mlp.jl)).

Is there any way to compute the confusion matrix in order to evaluate the performance of the model (need precision, recall and F1 score).

Thanks!

---

<div class="post-metadata">

**Author:** ![zgornel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zgornel/32/217487_2.png) [@zgornel](https://discourse.julialang.org/u/zgornel)\
**Post date:** [January 17, 2019, 10:38am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/2 "2019-01-17T10:38:05Z")

</div>

As a quick hack, you can find a simple implementation of a confusion matrix here: [j4pr.jl/libutils.jl at master · OxoaResearch/j4pr.jl · GitHub](https://github.com/OxoaResearch/j4pr.jl/blob/master/src/lib/libutils.jl) (may need a bit of tweaking for Julia \>0.6 …)

---

<div class="post-metadata">

**Author:** ![essenciary](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/essenciary/32/210469_2.png) [@essenciary](https://discourse.julialang.org/u/essenciary)\
**Post date:** [January 21, 2019, 2:39pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/3 "2019-01-21T14:39:58Z")

</div>

Thanks! In the end, I implemented it myself.  
I also discovered that there’s a `confusmat` function in MLBase but I’m not sure if/how it works with Flux.

---

<div class="post-metadata">

**Author:** ![zgornel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zgornel/32/217487_2.png) [@zgornel](https://discourse.julialang.org/u/zgornel)\
**Post date:** [January 21, 2019, 3:22pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/4 "2019-01-21T15:22:28Z")

</div>

Should work as well. The basic confusion matrix needs just two vectors ( references and predictions). Btw, I believe `MLBase` is informally deprecated in favour of `LearnBase`…

---

<div class="post-metadata">

**Author:** ![essenciary](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/essenciary/32/210469_2.png) [@essenciary](https://discourse.julialang.org/u/essenciary)\
**Post date:** [January 21, 2019, 7:05pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/5 "2019-01-21T19:05:31Z")

</div>

Oh, I see - makes sense as MLBase wasn’t updated in the last few months.

I’m not sure about using it with Flux - how would it work for example in this model: [https://github.com/FluxML/model-zoo/blob/master/vision/mnist/conv.jl](https://github.com/FluxML/model-zoo/blob/master/vision/mnist/conv.jl)

---

<div class="post-metadata">

**Author:** ![zgornel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zgornel/32/217487_2.png) [@zgornel](https://discourse.julialang.org/u/zgornel)\
**Post date:** [January 21, 2019, 9:22pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/6 "2019-01-21T21:22:27Z")

</div>

The model you pointed to is just the training. Applying it i.e. inference, should output a vector of 10 posterior probabilities for each sample (10 classes, 0 to 9). From there, the predicted label can be extracted (max posterior probability); the predicted labels can then be fed into the confusion matrix …

---

<div class="post-metadata">

**Author:** ![essenciary](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/essenciary/32/210469_2.png) [@essenciary](https://discourse.julialang.org/u/essenciary)\
**Post date:** [January 21, 2019, 10:09pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/7 "2019-01-21T22:09:39Z")

</div>

Thanks. It seems that some testing is performed throughout the training process, isn’t it? Instead of accuracy, it would be useful to be able to plug-in a full confusion matrix computation:

```julia
# Prepare test set (first 1,000 images)
tX = cat(float.(MNIST.images(:test)[1:1000])..., dims = 4) |> gpu
tY = onehotbatch(MNIST.labels(:test)[1:1000], 0:9) |> gpu

...

accuracy(x, y) = mean(onecold(m(x)) .== onecold(y))
evalcb = throttle(() -> @show(accuracy(tX, tY)), 10)

```

---

<div class="post-metadata">

**Author:** ![zgornel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zgornel/32/217487_2.png) [@zgornel](https://discourse.julialang.org/u/zgornel)\
**Post date:** [January 21, 2019, 10:20pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/8 "2019-01-21T22:20:44Z")

</div>

Couldn’t say, I’m really not familiar with Flux. Neural models backpropagate errors at sample level, not class level. In practice, they do optimize the confusion matrix by trying to learn the correct classes 🙂

---

<div class="post-metadata">

**Author:** ![Iulian.Cioarca](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/iulian.cioarca/32/30166_2.png) [@Iulian.Cioarca](https://discourse.julialang.org/u/Iulian.Cioarca)\
**Post date:** [February 4, 2020, 8:35am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/9 "2020-02-04T08:35:30Z")

</div>

Sorry for reviving this, is `MLBase` still maintained? I noticed there were some discussions on the github page of merging it into `StatsBase`.  
I was also searching for a confusion matrix implementation.

---

<div class="post-metadata">

**Author:** ![zgornel](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/zgornel/32/217487_2.png) [@zgornel](https://discourse.julialang.org/u/zgornel)\
**Post date:** [February 4, 2020, 10:05am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/10 "2020-02-04T10:05:59Z")

</div>

I guess so as it is part of `JuliaStats` yet it receives very little attention. I would not really rely on it.

---

<div class="post-metadata">

**Author:** ![Erro\_Ashivudhi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/erro_ashivudhi/32/24797_2.png) [@Erro\_Ashivudhi](https://discourse.julialang.org/u/Erro_Ashivudhi)\
**Post date:** [June 15, 2021, 2:47pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/11 "2021-06-15T14:47:19Z")

</div>

How did you implement it?

---

<div class="post-metadata">

**Author:** ![essenciary](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/essenciary/32/210469_2.png) [@essenciary](https://discourse.julialang.org/u/essenciary)\
**Post date:** [June 15, 2021, 3:09pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/12 "2021-06-15T15:09:02Z")

</div>

Can’t really remember, it was for a paper for my MSc - I think I’ve done it by hand (wrote the code from scratch).

---

<div class="post-metadata">

**Author:** ![Erro\_Ashivudhi](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/erro_ashivudhi/32/24797_2.png) [@Erro\_Ashivudhi](https://discourse.julialang.org/u/Erro_Ashivudhi)\
**Post date:** [June 15, 2021, 8:18pm UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/13 "2021-06-15T20:18:51Z")

</div>

Alright, thanks

---

<div class="post-metadata">

**Author:** ![TheLateKronos](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/thelatekronos/32/12824_2.png) [@TheLateKronos](https://discourse.julialang.org/u/TheLateKronos)\
**Post date:** [June 9, 2022, 9:19am UTC](https://discourse.julialang.org/t/flux-jl-confusion-matrix/19740/14 "2022-06-09T09:19:10Z")

</div>

See [Performance Measures · MLJ (alan-turing-institute.github.io)](https://alan-turing-institute.github.io/MLJ.jl/stable/performance_measures/)

```julia
julia> using CategoricalArrays, MLJ

julia> yhat = rand(1:10, 100)|>CategoricalArray;

julia> y = rand(1:10, 100)|>CategoricalArray;

julia> ConfusionMatrix()(yhat, y)
┌ Warning: The classes are un-ordered,
│ using order: [1, 2, 3, 4, 5, 6, 7, 8, 9, 10].
│ To suppress this warning, consider coercing to OrderedFactor.
└ @ MLJBase C:\Users\usrname\.julia\packages\MLJBase\rQDaq\src\measures\confusion_matrix.jl:122
10×10 Matrix{Int64}:
 1 3 1 3 1 0 1 0 1 1
 1 0 2 0 1 3 0 1 0 0
 1 2 1 0 3 1 0 1 1 1
 0 1 0 1 0 1 2 1 1 4
 1 1 1 1 3 1 2 0 3 3
 1 0 2 1 2 0 1 0 0 1
 0 1 0 1 1 2 1 1 0 1
 1 1 1 1 1 1 2 0 0 1
 2 0 0 2 1 1 1 1 0 0
 3 3 1 0 0 0 0 0 2 0

```
