# Instability with union of types with ForwardDiff and KeyedArray types, hints

**URL:** <https://discourse.julialang.org/t/instability-with-union-of-types-with-forwarddiff-and-keyedarray-types-hints/100772>\
**Category:** General Usage\
**Tags:** question, scope, code\_warntype, type-stability, setindex\
**Created:** [June 24, 2023, 4:51am UTC](https://discourse.julialang.org/t/instability-with-union-of-types-with-forwarddiff-and-keyedarray-types-hints/100772 "2023-06-24T04:51:45Z")\
**Posts on this page:** 2\
**Page:** 1

<div class="post-metadata">

**Author:** ![lazarusA](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/lazarusa/32/6571_2.png) [@lazarusA](https://discourse.julialang.org/u/lazarusA)\
**Post date:** [June 24, 2023, 4:51am UTC](https://discourse.julialang.org/t/instability-with-union-of-types-with-forwarddiff-and-keyedarray-types-hints/100772/1 "2023-06-24T04:51:45Z")

</div>

The goal is to calculate a number between two Vector-ish things in a type stable manner, without type assertions. This is part of a larger workflow that involves `ForwardDiff.gradient`, which runs, but slow, hence the inquiry 😃

```plaintext
using Dates, ForwardDiff, AxisKeys
r = rand(Float32, 4)
r[1] = NaN32    
ka = KeyedArray(r; i =Date(today()):Date(today()+Day(3)) )
ar = Union{Float32, ForwardDiff.Dual}[1f0,2f0, ForwardDiff.Dual(NaN32), 2f0]
idx_f(ar,ka) = (.!isnan.(ar .* ka))
idxs = idx_f(ar,ka)

```

For this initial function we already see some instabilities

> **@code\_warntype idx\_f(ar,ka)**
>
> Arguments #self#::Core.Const(idx\_f) ar::Vector{Union{Float32, ForwardDiff.Dual}} ka::KeyedArray{Float32, 1, NamedDimsArray{(:i,), Float32, 1, Vector{Float32}}, Base.RefValue{StepRange{Date, Day}}} Body::KeyedArray{\_A, 1, \_B, Base.RefValue{StepRange{Date, Day}}} where {\_A, \_B} 1 ─ %1 = Main.:!::Core.Const(!) │ %2 = Main.isnan::Core.Const(isnan) │ %3 = Base.broadcasted(Main.:\*, ar, ka)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(\*), T} │ %4 = Base.broadcasted(%2, %3)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(isnan), NT} │ %5 = Base.broadcasted(%1, %4)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(!), NT} │ %6 = Base.materialize(%5)::KeyedArray{\_A, 1, \_B, Base.RefValue{StepRange{Date, Day}}} where {\_A, \_B} └── return %6

and let’s say that the target function is

```plaintext
function ka_ar_unstable(ka, ar,idxs)
    return abs2.(ka[idxs] .- ar[idxs])
end

```

> **@code\_warntype ka\_ar\_unstable(ka, ar, idxs)**
>
> Arguments #self#::Core.Const(ka\_ar\_unstable) ka::KeyedArray{Float32, 1, NamedDimsArray{(:i,), Float32, 1, Vector{Float32}}, Base.RefValue{StepRange{Date, Day}}} ar::Vector{Union{Float32, ForwardDiff.Dual}} idxs::KeyedArray{Bool, 1, NamedDimsArray{(:i,), Bool, 1, BitVector}, Base.RefValue{StepRange{Date, Day}}} Body::KeyedArray{\_A, 1, \_B, Base.RefValue{Vector{Date}}} where {\_A, \_B} 1 ─ %1 = Main.abs2::Core.Const(abs2) │ %2 = Main.:-::Core.Const(-) │ %3 = Base.getindex(ka, idxs)::KeyedArray{Float32, 1, NamedDimsArray{(:i,), Float32, 1, Vector{Float32}}, Base.RefValue{Vector{Date}}} │ %4 = Base.getindex(ar, idxs)::Vector{Union{Float32, ForwardDiff.Dual}} │ %5 = Base.broadcasted(%2, %3, %4)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(-), T} │ %6 = Base.broadcasted(%1, %5)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(abs2), NT} │ %7 = Base.materialize(%6)::KeyedArray{\_A, 1, \_B, Base.RefValue{Vector{Date}}} where {\_A, \_B} └── return %7

and all togheter

```julia
function all_pack(ar, ka)
    idts = idx_f(ar,ka) # this needs to be called here because, 'ar' and 'ka' 
                        #are comming from an outer loop.
    vals = ka_ar_unstable(ka, ar,idxs)
    return sum(vals)
end

```

> **@code\_warntype all\_pack(ar, ka)**
>
> Arguments #self#::Core.Const(all\_pack) ar::Vector{Union{Float32, ForwardDiff.Dual}} ka::KeyedArray{Float32, 1, NamedDimsArray{(:i,), Float32, 1, Vector{Float32}}, Base.RefValue{StepRange{Date, Day}}} Locals vals::Any idxs::KeyedArray{\_A, 1, \_B, Base.RefValue{StepRange{Date, Day}}} where {\_A, \_B} Body::Any 1 ─ (idxs = Main.idx\_f(ar, ka)) │ (vals = Main.ka\_ar\_unstable(ka, ar, idxs)) │ %3 = Main.sum(vals)::Any └── return %3

doing the indices inside leads to `Base.getindex` issues directly.

```plaintext
function ka_ar_unstable(ka, ar)
    idxs = (.!isnan.(ar .* ka))
    return abs2.(ka[idxs] .- ar[idxs])
end

```

> **@code\_warntype ka\_ar\_unstable(ar, ka)**
>
> Arguments  
> #self#::Core.Const(ka\_ar\_unstable)  
> ka::Vector{Union{Float32, ForwardDiff.Dual}}  
> ar::KeyedArray{Float32, 1, NamedDimsArray{(:i,), Float32, 1, Vector{Float32}}, Base.RefValue{StepRange{Date, Day}}}  
> Locals  
> idxs::KeyedArray{\_A, 1, \_B, Base.RefValue{StepRange{Date, Day}}} where {\_A, \_B}  
> Body::Any  
> 1 ─ %1 = Main.:!::Core.Const(!)  
> │ %2 = Main.isnan::Core.Const(isnan)  
> │ %3 = Base.broadcasted(Main.:_, ar, ka)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(_), T}  
> │ %4 = Base.broadcasted(%2, %3)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(isnan), NT}  
> │ %5 = Base.broadcasted(%1, %4)::Base.Broadcast.Broadcasted{AxisKeys.KeyedStyle{NamedDims.NamedDimsStyle{Base.Broadcast.DefaultArrayStyle{1}}}, Nothing, typeof(!), NT}  
> │ (idxs = Base.materialize(%5))  
> │ %7 = Main.abs2::Core.Const(abs2)  
> │ %8 = Main.:-::Core.Const(-)  
> │ %9 = Base.getindex(ka, idxs)::Any  
> │ %10 = Base.getindex(ar, idxs)::Any  
> │ %11 = Base.broadcasted(%8, %9, %10)::Any  
> │ %12 = Base.broadcasted(%7, %11)::Any  
> │ %13 = Base.materialize(%12)::Any  
> └── return %13

any ideas or hints that could help solve this issue would be greatly appreciated 😃 . Maybe is something very trivial for some 🙂 .

```julia
pkg> status Dates ForwardDiff AxisKeys
Status `~/Project.toml`
  [94b1ba4f] AxisKeys v0.2.13
  [f6369f11] ForwardDiff v0.10.35

```

---

<div class="post-metadata">

**Author:** ![fabiangans](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/fabiangans/32/2624_2.png) [@fabiangans](https://discourse.julialang.org/u/fabiangans)\
**Post date:** [June 27, 2023, 8:07am UTC](https://discourse.julialang.org/t/instability-with-union-of-types-with-forwarddiff-and-keyedarray-types-hints/100772/2 "2023-06-27T08:07:08Z")

</div>

I think the main reason for the instability is that `ForwardDiff.Dual` is not a concrete type, so your pre-allocated array `ar` can per-se not be inferred. Note that the type-instability is gone when you just define your array as

```julia
ar = [1f0,2f0, ForwardDiff.Dual(NaN32), 2f0]

```

However, I guess you need to pre-allocate the array so it can be re-used. I think an option would be to use PreallocationsTools.jl as in this example [GitHub - SciML/PreallocationTools.jl: Tools for building non-allocating pre-cached functions in Julia, allowing for GC-free usage of automatic differentiation in complex codes](https://github.com/SciML/PreallocationTools.jl#diffcache-example-1-direct-usage) .
