# Flux Dense Layer Type Instability

**URL:** <https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642>\
**Category:** Machine Learning\
**Created:** [May 10, 2023, 8:49pm UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642 "2023-05-10T20:49:35Z")\
**Posts on this page:** 8\
**Page:** 1

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 10, 2023, 8:49pm UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/1 "2023-05-10T20:49:35Z")

</div>

Hi. I am stumped as to why the forward pass is type unstable. It’s effectively a varying length chain similar to Flux.Chain on an alternating sequence of dense and dropout layers:

```julia
struct MLP
    dense::Vector{Flux.Dense}
    drop::Vector{Flux.Dropout}
end

function MLP(layer_dims::Vector{Int}, dropout::Bool=true, activation=tanh)
    dense = Flux.Dense[]
    drop = Flux.Dropout[]
    for i=1:length(layer_dims)-1
        if dropout && i > 1
            push!(drop, Flux.Dropout(0.5))
        end

        if i < length(layer_dims)
            push!(dense, Flux.Dense(layer_dims[i]=>layer_dims[i+1], activation))
        else
            push!(dense, Flux.Dense(layer[i] => layer[i+1]))
        end
    end
end

function (mlp::MLP)(x::Matrix{Float32}) # Forward pass
    temp = x
    for i=1:length(mlp.drop)  
        temp = mlp.dense[i](temp)
        temp = mlp.drop[i](temp)
    end
    mlp.dense[i](temp)
    mlp.dense[end](temp)
    temp
end

mlp = SR.MLP([2,2,])
x_fake = rand(Float32, 2, 100)
@code_warntype mlp(x_fake)

```

The lowered rep shows that `temp` is unstable:

```julia
MethodInstance for ()
Arguments
  mlp::MLP
  x::Matrix{Float32}
Locals
  @_3::Union{Nothing, Tuple{Int64, Int64}}
  temp::Any 
  i::Int64
Body::Any
1 ─ (temp = x)
│ %2 = Base.getproperty(mlp, :drop)::Vector{Flux.Dropout}
│ %3 = length(%2)::Int64
│ %4 = (1:%3)::Core.PartialStruct(UnitRange{Int64}, Any[Core.Const(1), Int64])
│ (@_3 = Base.iterate(%4))
│ %6 = (@_3 === nothing)::Bool
│ %7 = Base.not_int(%6)::Bool
└── goto #4 if not %7
2 ┄ %9 = @_3::Tuple{Int64, Int64}
│ (i = Core.getfield(%9, 1))
│ %11 = Core.getfield(%9, 2)::Int64
│ %12 = Base.getproperty(mlp, :dense)::Vector{Flux.Dense}
│ %13 = Base.getindex(%12, i)::Flux.Dense
│ (temp = (%13)(temp))
│ %15 = Base.getproperty(mlp, :drop)::Vector{Flux.Dropout}
│ %16 = Base.getindex(%15, i)::Flux.Dropout
│ (temp = (%16)(temp))
│ (@_3 = Base.iterate(%4, %11))
│ %19 = (@_3 === nothing)::Bool
│ %20 = Base.not_int(%19)::Bool
└── goto #4 if not %20
3 ─ goto #2
4 ┄ %23 = Base.getproperty(mlp, :dense)::Vector{Flux.Dense}
│ %24 = Base.getindex(%23, i)::Any
│ (%24)(temp)
│ %26 = Base.getproperty(mlp, :dense)::Vector{Flux.Dense}
│ %27 = Base.lastindex(%26)::Int64
│ %28 = Base.getindex(%26, %27)::Flux.Dense
│ (%28)(temp)
└── return temp

```

I even tried evaluating with one Dense layer and the output it still Any. Is there anything I can do about this or is it not a problem?

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [May 11, 2023, 5:41am UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/2 "2023-05-11T05:41:18Z")

</div>

Why do you need an alternative to `Flux.Chain`? I think your problem comes from your custom struct not being concretely typed, as you can check by running `isconcretetype(typeof(mlp))`. When you take a look at structs like `Flux.Dense`, they have type parameters to specify what’s inside, which are left aside in your vector storage:

> <https://github.com/FluxML/Flux.jl/blob/c9c262db1c851cc612389f86854b1987083aab25/src/layers/basic.jl#L153-L156>

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 11, 2023, 6:40am UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/3 "2023-05-11T06:40:29Z")

</div>

So this mlp is a submodel of a larger struct which calls the mlp at the end of its forward pass. I originally tried using `Chain`, but the problem is that for some reason I received similar warnings from @code\_warntype.

I should have put this in the OP, but here is an MWE of the bigger models forward pass:

```julia
function (model::BiggerModel)(x, λ, ϕ, A, ∇_x, ∇_y)
    x_diffused = model.diffusion_block(x, λ, ϕ, A)
    x_intermediate = vcat(x_diffused', x_intermediate)
    x_out = model.mlp(x_intermediate) # Any!
end

```

This happens if `mlp` is either the implementation above, or `Flux.Chains`. Here is also the code I used to construct the `Chain`:

```julia
function MLP(layer_dims::Vector{Int}, dropout::Bool=false, activation=tanh)
    layers = Union{Flux.Dropout, Flux.Dense}[]
    for i=1:length(layer_dims)-1
        if dropout && i > 0
            push!(layers, Flux.Dropout(0.5))
        end
        if i < length(layer_dims)
            push!(layers, Flux.Dense(layer_dims[i] => layer_dims[i+1],activation))
        else
            push!(layers, Flux.Dense(layer[i]=>layer[i+1]))
        end
    end
    Flux.Chain(layers...)
end

```

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 11, 2023, 6:44am UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/4 "2023-05-11T06:44:40Z")

</div>

And here is an even simpler example what of I’m concerned about:

```julia
let
	struct Foo
		d::Flux.Chain
	end
	@Flux.functor Foo
	function (model::Foo)(x)
		temp = d(x)
	end
	d = Flux.Chain(Dense(2=>2), Dropout(0.5))
	g = Foo(d)
	@code_warntype g(rand(Float32, (2,10)))
end

```

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [May 11, 2023, 6:52am UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/5 "2023-05-11T06:52:08Z")

</div>

Can you give the struct definition for `BiggerModel`?

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [May 11, 2023, 2:06pm UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/6 "2023-05-11T14:06:38Z")

</div>

As Guillaume said, declaring types which contain `Dense` or `Chain` requires you to provide the type parameters for those structs as well somewhere in the wrapping type (e.g. `MLP`). Otherwise the types aren’t fully specified and everything is type unstable. See [Performance Tips · The Julia Language](https://docs.julialang.org/en/v1/manual/performance-tips/#Avoid-fields-with-abstract-containers) in the performance tips for more info.

---

<div class="post-metadata">

**Author:** ![Ian\_L](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/ian_l/32/49509_2.png) [@Ian\_L](https://discourse.julialang.org/u/Ian_L)\
**Post date:** [May 11, 2023, 4:10pm UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/7 "2023-05-11T16:10:50Z")

</div>

> [@ToucheSir](#):
>
> requires you to provide the type parameters for those structs as well somewhere in the wrapping type (e.g. `MLP` ). Otherwise the types aren’t fully specified and everything is type unstable.

Ok. So if I would want to make everything concrete, I would need to also need to parameterize `MLP` with the same parameters as dense? This works:

```julia
	struct Foo{F,M<:AbstractMatrix,B}
		d::Flux.Dense{F, M, B}
	end
	f = Foo(Dense(2=>2))
	function (f::Foo)(x)
		temp = f.d(x)
	end
	@code_warntype f(rand(Float32, 2,10))

```

But would mean that if Foo is part of another larger model (in my case there is another one), then I would also need to parameterize the larger model as well? Is there a more concise way to define Foo?

---

<div class="post-metadata">

**Author:** ![ToucheSir](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/touchesir/32/14411_2.png) [@ToucheSir](https://discourse.julialang.org/u/ToucheSir)\
**Post date:** [May 11, 2023, 4:20pm UTC](https://discourse.julialang.org/t/flux-dense-layer-type-instability/98642/8 "2023-05-11T16:20:12Z")

</div>

```julia
struct Foo{D<:Dense}
  d::D
end

```

Or even remove the type constraint entirely:

```julia
struct Foo{L}
  layer::L
end

```

Which is what most of Flux’s container layers do.
