# How to use crossentropy in a multi-class classification Flux.jl neural network?

**URL:** <https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140>\
**Category:** Machine Learning\
**Tags:** first-steps, flux\
**Created:** [October 10, 2018, 10:10pm UTC](https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140 "2018-10-10T22:10:59Z")\
**Posts on this page:** 4\
**Page:** 1

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [October 10, 2018, 10:10pm UTC](https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140/1 "2018-10-10T22:10:59Z")

</div>

I am having trouble defining the `crossentropy` loss using Flux.jl.

```julia
using Flux,StatsBase
model = Flux.Chain(
  Dense(13*16, 128, relu),
  Dense(128, 64, relu),
  Dense(64, 32, relu),
  Dense(32, 4, relu),
    softmax);

loss(x,y) = crossentropy(model(x),y)
opt = ADAM(params(model))

```

I have set up the model to try and predict a 4-label classification problem, but I can’t seem to get the `loss` function to work. What form does `y` have to be?

For example, my `y` for a record can be coded as `[0.0,1.0,0.0,0.0]` but running `crossentropy` gives `Inf (Tracked)`.

---

<div class="post-metadata">

**Author:** ![Tomas\_Pevny](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/tomas_pevny/32/25466_2.png) [@Tomas\_Pevny](https://discourse.julialang.org/u/Tomas_Pevny)\
**Post date:** [October 11, 2018, 4:55am UTC](https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140/2 "2018-10-11T04:55:04Z")

</div>

Flux expect `y` to be matrix. `crossentropy(model(x),onehotbatch(y, 0:1))` should work, assuming that `y` is and Int vector of labels with 0 and 1s.

`onehotbatch` is a function defined in Flux

---

<div class="post-metadata">

**Author:** ![baggepinnen](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/baggepinnen/32/693_2.png) [@baggepinnen](https://discourse.julialang.org/u/baggepinnen)\
**Post date:** [October 11, 2018, 6:23am UTC](https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140/3 "2018-10-11T06:23:41Z")

</div>

The [model-zoo](https://github.com/FluxML/model-zoo) contains [many example models using crossentropy](https://github.com/FluxML/model-zoo/search?q=crossentropy&unscoped_q=crossentropy). You can use them to see how common models and patterns are implemented in Flux

---

<div class="post-metadata">

**Author:** ![xiaodai](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/xiaodai/32/15937_2.png) [@xiaodai](https://discourse.julialang.org/u/xiaodai)\
**Post date:** [October 11, 2018, 10:18am UTC](https://discourse.julialang.org/t/how-to-use-crossentropy-in-a-multi-class-classification-flux-jl-neural-network/16140/4 "2018-10-11T10:18:37Z")

</div>

Both `Flux.jl` and `StatsBase` exports `crossentropy`, and using `loss(x,y) = Flux.crossentropy(model(x),y)` solves the problem
