# Solving type instabilities in simple classification model

**URL:** <https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572>\
**Category:** New to Julia\
**Created:** [October 30, 2023, 2:21pm UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572 "2023-10-30T14:21:55Z")\
**Posts on this page:** 8\
**Page:** 1

<div class="post-metadata">

**Author:** ![looking4stability](https://avatars.discourse-cdn.com/v4/letter/l/d07c76/32.png) [@looking4stability](https://discourse.julialang.org/u/looking4stability)\
**Post date:** [October 30, 2023, 2:21pm UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/1 "2023-10-30T14:21:55Z")

</div>

Hi,  
I’m pretty new to Julia, so please be nice:)

I wrote a simple script to identify mnist digits.  
So far it runs okay, and I like the language a lot.  
But what I can’t get my head wrapped around are the type instabilities I’m trying to solve.

When I am calling this function:

```julia
function create_A_model()
    Chain(
        Dense(784, 512, relu), 
        Dense(512, 10),
        softmax)         
end

```

This is the output from code\_warntype:

Arguments  
#self#::Core.Const(create\_model)  
config::TrainingConfig  
Body **::Chain**  
1 ─ %1 = Base.getproperty(config, :input\_size)::Int64  
│ %2 = Base.getproperty(config, :hidden\_size)::Int64  
│ %3 = Main.Dense(%1, %2, Main.relu) **::Dense{typeof(relu), Matrix{Float32}}**  
│ %4 = Base.getproperty(config, :hidden\_size)::Int64  
│ %5 = Base.getproperty(config, :output\_size)::Int64  
│ %6 = Main.Dense(%4, %5)**::Dense{typeof(identity), Matrix{Float32}}**  
│ %7 = Main.Chain(%3, %6, Main.softmax) **::Chain**  
└── return %7  
I marked the red parts in bold

What’s even more confusing to me is that if I reduce the function to

```julia
function create_model()
    Chain(Dense(1,1)) 
end

```

I can only post 1 pictur, so I removed the output here, but  
::Chain  
::Dense{typeof(identity), Matrix{Float32}}  
is still red

I have tried to use structs and I can’t seem to find answers in the documentation.

The other type instabilities I can’t get rid of, are within my training function.  
Originally I had more functionality, but I stopped it down to pinpoint the problem.

```julia
function train_network(
    model::Chain{Tuple{Dense{typeof(relu), Matrix{Float32}, Vector{Float32}}, Dense{typeof(identity), Matrix{Float32}, Vector{Float32}}, typeof(softmax)}}, 
    train_data::DataLoader{Tuple{Matrix{Float32}, Flux.OneHotMatrix{UInt32, Vector{UInt32}}}}, 
    opt::Descent)
    epochstest::Int64 = 5
    for epoch::Int64 in 1:epochstest::Int64
        for (input::Matrix{Float32}, labels::Flux.OneHotMatrix{UInt32, Vector{UInt32}}) in train_data
            Flux.train!((input, labels) -> loss_function(input, labels, model), Flux.params(model), [(input, labels)], opt)
        end
    end
end

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/8/0/80eec6056c526335f6283a0a69798f9ec0c5d10a.png)

All other functions in my code are type-stable according to code\_warntype.  
when I understand the output correctly, then \_5 has partial instability.  
Are \_5, \_7 and \_10 internal elements? How could I make them type-stable?  
Or are they unstable because the Chain in the earlier function is unstable?

I would be really glad for any kind of helpful input, because I’m not getting closer to a solution no matter what I try and I’ve been trying for too long 😅

if I can provide any more information or more code, I’m happy to do so

thanks  
l4s

---

<div class="post-metadata">

**Author:** ![ParadaCarleton](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/paradacarleton/32/20005_2.png) [@ParadaCarleton](https://discourse.julialang.org/u/ParadaCarleton)\
**Post date:** [November 3, 2023, 2:14am UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/2 "2023-11-03T02:14:15Z")

</div>

First thing’s first, could you try removing all the type annotations from your code? It looks way overtyped and hard to read 😅

Julia can infer types. Usually, it does a better job than you can do annotating them them, so we usually only provide the bare minimum typing to make sure all of our inputs are correct. Adding more types doesn’t make the code any faster, and is a frequent source of bugs. Here’s one way to rewrite this:

```julia
function train_network(model::Chain, train_data::DataLoader, opt::Optimiser)

    epochstest = 5
    for epoch in 1:epochstest
        for (input, labels) in train_data
            Flux.train!(
                (input, labels) -> loss_function(input, labels, model), 
                Flux.params(model), 
                [(input, labels)], 
                opt
            )
        end
    end
end

```

Now that the code is cleaner and more readable, can you spot any mistakes, and is it there still a type instability?

---

<div class="post-metadata">

**Author:** ![abraemer](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/abraemer/32/51403_2.png) [@abraemer](https://discourse.julialang.org/u/abraemer)\
**Post date:** [November 3, 2023, 7:53am UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/3 "2023-11-03T07:53:19Z")

</div>

Welcome to the discourse!

Is there any particular reason you strive for type-stability in these function? Usually you want type-stability if you need maximal performance in a piece of code but neither of your function seem to me to be bottlenecks. You probably generate a model like once in while and `Flux.train!` probably does the heavy lifting for you.

Note that type-stability/inferrability is not contagious in any form. If your `train_network` function is not fully inferrable, then Julia needs to do dynamic dispatch for every call with inputs where Julia couldn’t figure out the precise types. But after then Julia dispatches to the method with precisely known types, so the type inference game starts anew! Probably the authors of Flux.jl took care and on their side things are inferrable and so you get good performance.

In fact it is quite often the case that the “entrance layer” to a library contains type-unstable code (for example in DifferentialEquations.jl to figure out which solvers to use) but beyond that type-unstable setup layer everything is inferrable and so you get maximal flexibility and maximal performance. This pattern is called [“function barrier” in the performance tips](https://docs.julialang.org/en/v1/manual/performance-tips/#kernel-functions).

---

<div class="post-metadata">

**Author:** ![looking4stability](https://avatars.discourse-cdn.com/v4/letter/l/d07c76/32.png) [@looking4stability](https://discourse.julialang.org/u/looking4stability)\
**Post date:** [November 4, 2023, 10:47am UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/4 "2023-11-04T10:47:48Z")

</div>

Thank you for your answers!

The overtyping was my attempt to solve the type-instabilities that the @code\_warntype macro warned me about.  
In the meantime I have read a lot about type instabilities in Julia and that their impact on performance is not always that big.

Furthermore I found some other means to increase the performance of my code and I’m quite happy now, that the training of the network only takes about 1/3 of the time that an equivalent PyTorch implementation needs.

This was the reason why I hoped to get a performance boost out of achieving perfect type-stability anyways.

I marked your answer as the solution @abraemer

---

<div class="post-metadata">

**Author:** ![ParadaCarleton](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/paradacarleton/32/20005_2.png) [@ParadaCarleton](https://discourse.julialang.org/u/ParadaCarleton)\
**Post date:** [November 4, 2023, 5:49pm UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/5 "2023-11-04T17:49:51Z")

</div>

> [@looking4stability](#):
>
> Furthermore I found some other means to increase the performance of my code and I’m quite happy now, that the training of the network only takes about 1/3 of the time that an equivalent PyTorch implementation needs.

Curious here, are you comparing to basic Pytorch or `@torch.compile`?

---

<div class="post-metadata">

**Author:** ![looking4stability](https://avatars.discourse-cdn.com/v4/letter/l/d07c76/32.png) [@looking4stability](https://discourse.julialang.org/u/looking4stability)\
**Post date:** [November 5, 2023, 10:50am UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/6 "2023-11-05T10:50:22Z")

</div>

I am comparing basic Pytorch to Flux to Basic Julia (no advanced packages).

I played around with torch.compile, but it didn’t produce performance increases. To my understanding the effect is most noticeable when using a GPU, having complex models as well as exporting the completed model to a different programming environment.

I’m using CPU only, the model is very basic…

---

<div class="post-metadata">

**Author:** ![ParadaCarleton](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/paradacarleton/32/20005_2.png) [@ParadaCarleton](https://discourse.julialang.org/u/ParadaCarleton)\
**Post date:** [November 5, 2023, 6:55pm UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/7 "2023-11-05T18:55:21Z")

</div>

> [@looking4stability](#):
>
> I played around with torch.compile, but it didn’t produce performance increases.

Interesting, did you try asking for help on the PyTorch forums?

---

<div class="post-metadata">

**Author:** ![looking4stability](https://avatars.discourse-cdn.com/v4/letter/l/d07c76/32.png) [@looking4stability](https://discourse.julialang.org/u/looking4stability)\
**Post date:** [November 7, 2023, 9:21am UTC](https://discourse.julialang.org/t/solving-type-instabilities-in-simple-classification-model/105572/8 "2023-11-07T09:21:03Z")

</div>

No I haven’t since it’s out of the scope of the study at hand, but it definitely is an interesting development for python.
