# Flux chain type unstable when broadcasting inside gradient

**URL:** <https://discourse.julialang.org/t/flux-chain-type-unstable-when-broadcasting-inside-gradient/115741>\
**Category:** Specific Domains\
**Tags:** flux, zygote, autodiff\
**Created:** [June 17, 2024, 9:27am UTC](https://discourse.julialang.org/t/flux-chain-type-unstable-when-broadcasting-inside-gradient/115741 "2024-06-17T09:27:22Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![filchristou](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/filchristou/32/26760_2.png) [@filchristou](https://discourse.julialang.org/u/filchristou)\
**Post date:** [June 17, 2024, 9:27am UTC](https://discourse.julialang.org/t/flux-chain-type-unstable-when-broadcasting-inside-gradient/115741/1 "2024-06-17T09:27:22Z")

</div>

I’ve been fighting for a couple of days to get my `Flux`/`Zygote` autodiff code type stable.  
I am not sure what’s the problem, but coming up with a MWE, it looks like broadcasting `Flux.Chain` is problematic ?

```julia
using Flux, Zygote
import Statistics: mean

function internfunc_nobroad(m, x, y)
    modelvals = m(x)
    Flux.mse(modelvals, y)
end

function internfunc_broad(m, x, y)
    modelvals = m.(x)
    mses = Flux.mse.(modelvals, y)
    return mean(mses)
end

function wrapfunc(model, xdata, ydata, func)
    grad = let xdata=xdata, ydata=ydata
        Zygote.gradient(m -> func(m, xdata, ydata), model)
    end
    return grad
end

```

Run the following in REPL

```julia
julia> fc = Flux.Chain(Flux.Dense(5=>3, Flux.relu), Flux.Dense(3=>3, Flux.relu), Flux.Dense(3=>1))
julia> fx = [fill(5f0, 5) for _ in 1:10]
julia> fy = fill(2f0, 10)

```

```julia
julia> @code_warntype wrapfunc(fc, fx, fy, internfunc_broad) # type unstable

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/1/6/16cf908866dc8e370a5d2c9aa4e6d400db853243.png)

```julia
julia> @code_warntype wrapfunc(fc, fx[1], fy[1], internfunc_nobroad) # type stable

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/0/5/0556b23918363e2cb737ca7a6b7268579242a587.png)

I made a similar [issue in Flux.jl](https://github.com/FluxML/Flux.jl/issues/2456)

---

<div class="post-metadata">

**Author:** ![filchristou](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/filchristou/32/26760_2.png) [@filchristou](https://discourse.julialang.org/u/filchristou)\
**Post date:** [June 17, 2024, 11:57am UTC](https://discourse.julialang.org/t/flux-chain-type-unstable-when-broadcasting-inside-gradient/115741/2 "2024-06-17T11:57:28Z")

</div>

Okey, I think I got it… I should convert the input to a matrix and not a Vector of Vectors. Then, Flux handles that nicely.

```julia
fobs_ar = fill(5f0, 5, 10)
labels_ar = fill(2f0, 1, 10)

@code_warntype wrapfunc(fc, fobs_ar, labels_ar, internfunc_nobroad)

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/6/d/6d9fa8bb4da9fd2a05add69b890859ed5f378638.png)

---

<div class="post-metadata">

**Author:** ![filchristou](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/filchristou/32/26760_2.png) [@filchristou](https://discourse.julialang.org/u/filchristou)\
**Post date:** [June 17, 2024, 2:19pm UTC](https://discourse.julialang.org/t/flux-chain-type-unstable-when-broadcasting-inside-gradient/115741/3 "2024-06-17T14:19:13Z")

</div>

well… After switching from `Flux.mse` to `Flux.huber_loss` I get type unstable code again…

```julia
function internfunc_nobroad_huberloss(m, x, y)
    modelvals = m(x)
    Flux.huber_loss(modelvals, y)
end

@code_warntype wrapfunc(fc, fobs_ar, labels_ar, internfunc_nobroad)

```

 ![image](https://global.discourse-cdn.com/julialang/original/3X/6/c/6c962306c79a8bc838973cfa20a2b159d2b1c153.png)

This looks definitely like a bug.  
I [made an issue](https://github.com/FluxML/Flux.jl/issues/2459). Feel free to drop some hints if you know why is that and how could it be tackled.
