# Hessian Vector Products on GPU using ForwardDiff, Zygote, and Flux

**URL:** <https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306>\
**Category:** Machine Learning\
**Tags:** gpu, cuda, flux, zygote, forwarddiff\
**Created:** [January 27, 2022, 4:48pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306 "2022-01-27T16:48:15Z")\
**Posts on this page:** 14\
**Page:** 1

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 27, 2022, 4:48pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/1 "2022-01-27T16:48:15Z")

</div>

Hello,

I am trying to compute the Hessian vector product of a Flux model (with some vector) on the GPU. This is giving me some errors, which I think are realted specifically to ForwardDiff. Any help would be greatly appreciated!

The following defines an Hvp operator which works fine on the CPU:

```julia
using ForwardDiff: partials, Dual
using Zygote: pullback
using LinearAlgebra

mutable struct HvpOperator{F, T, I}
	f::F
	x::AbstractArray{T, 1}
	dualCache1::AbstractArray{Dual{Nothing, T, 1}}
	size::I
	nProd::I
end

function HvpOperator(f, x::AbstractVector)
	dualCache1 = Dual.(x, similar(x))
	return HvpOperator(f, x, dualCache1, size(x, 1), 0)
end

Base.eltype(op::HvpOperator{F, T, I}) where{F, T, I} = T
Base.size(op::HvpOperator) = (op.size, op.size)

function LinearAlgebra.mul!(result::AbstractVector, op::HvpOperator, v::AbstractVector)
	op.nProd += 1

	op.dualCache1 .= Dual.(op.x, v)
	val, back = pullback(op.f, op.dualCache1)

	result .= partials.(back(one(val))[1], 1)
end

```

Ignore the `Nothing` tags for now as I know I shouldn’t be doing that. The following code defines a MWE to use the Hvp operator on the GPU:

```julia
using Flux, CUDA
data = randn(10, 4) |> gpu
model = Dense(10, 1, σ) |> gpu

ps, re = Flux.destructure(model)
f(θ) = re(θ)(data)

Hop = HvpOperator(f, ps)

v, res = similar(ps), similar(ps)

LinearAlgebra.mul!(res, Hop, v)

```

This results in the following error:

```julia
ERROR: LoadError: MethodError: no method matching cudnnDataType(::Type{Dual{Nothing, Float32, 1}})
Closest candidates are:
  cudnnDataType(::Type{Float16}) at /home/rs-coop/.julia/packages/CUDA/nYggH/lib/cudnn/util.jl:7
  cudnnDataType(::Type{Float32}) at /home/rs-coop/.julia/packages/CUDA/nYggH/lib/cudnn/util.jl:8
  cudnnDataType(::Type{Float64}) at /home/rs-coop/.julia/packages/CUDA/nYggH/lib/cudnn/util.jl:9
  ...
Stacktrace:
  [1] CUDA.CUDNN.cudnnTensorDescriptor(array::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}; format::CUDA.CUDNN.cudnnTensorFormat_t, dims::Vector{Int32})
    @ CUDA.CUDNN ~/.julia/packages/CUDA/nYggH/lib/cudnn/tensor.jl:9
  [2] CUDA.CUDNN.cudnnTensorDescriptor(array::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer})
    @ CUDA.CUDNN ~/.julia/packages/CUDA/nYggH/lib/cudnn/tensor.jl:8
  [3] cudnnActivationForward!(y::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}, x::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}; o::Base.Iterators.Pairs{Symbol, CUDA.CUDNN.cudnnActivationMode_t, Tuple{Symbol}, NamedTuple{(:mode,), Tuple{CUDA.CUDNN.cudnnActivationMode_t}}})
    @ CUDA.CUDNN ~/.julia/packages/CUDA/nYggH/lib/cudnn/activation.jl:22
  [4] (::NNlibCUDA.var"#65#69")(src::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}, dst::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer})
    @ NNlibCUDA ~/.julia/packages/NNlibCUDA/IeeBk/src/cudnn/activations.jl:11
  [5] materialize(bc::Base.Broadcast.Broadcasted{CUDA.CuArrayStyle{2}, Nothing, typeof(σ), Tuple{CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}}})
    @ NNlibCUDA ~/.julia/packages/NNlibCUDA/IeeBk/src/cudnn/activations.jl:30
  [6] rrule(#unused#::typeof(Base.Broadcast.broadcasted), #unused#::typeof(σ), x::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer})
    @ NNlib ~/.julia/packages/NNlib/y5z4i/src/activations.jl:813
  [7] rrule(::Zygote.ZygoteRuleConfig{Zygote.Context}, ::Function, ::Function, ::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer})
    @ ChainRulesCore ~/.julia/packages/ChainRulesCore/oBjCg/src/rules.jl:134
  [8] chain_rrule
    @ ~/.julia/packages/Zygote/FPUm3/src/compiler/chainrules.jl:216 [inlined]
  [9] macro expansion
    @ ~/.julia/packages/Zygote/FPUm3/src/compiler/interface2.jl:0 [inlined]
 [10] _pullback(::Zygote.Context, ::typeof(Base.Broadcast.broadcasted), ::typeof(σ), ::CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer})
    @ Zygote ~/.julia/packages/Zygote/FPUm3/src/compiler/interface2.jl:9
 [11] _pullback
    @ ~/.julia/packages/Flux/BPPNj/src/layers/basic.jl:158 [inlined]
 [12] _pullback(ctx::Zygote.Context, f::Dense{typeof(σ), CuArray{Dual{Nothing, Float32, 1}, 2, CUDA.Mem.DeviceBuffer}, CuArray{Dual{Nothing, Float32, 1}, 1, CUDA.Mem.DeviceBuffer}}, args::CuArray{Float32, 2, CUDA.Mem.DeviceBuffer})
    @ Zygote ~/.julia/packages/Zygote/FPUm3/src/compiler/interface2.jl:0
 [13] _pullback
    @ ~/Documents/research/masters-thesis/CubicNewton/experiments/test.jl:43 [inlined]
 [14] _pullback(ctx::Zygote.Context, f::typeof(f), args::CuArray{Dual{Nothing, Float32, 1}, 1, CUDA.Mem.DeviceBuffer})
    @ Zygote ~/.julia/packages/Zygote/FPUm3/src/compiler/interface2.jl:0
 [15] _pullback(f::Function, args::CuArray{Dual{Nothing, Float32, 1}, 1, CUDA.Mem.DeviceBuffer})
    @ Zygote ~/.julia/packages/Zygote/FPUm3/src/compiler/interface.jl:34
 [16] pullback(f::Function, args::CuArray{Dual{Nothing, Float32, 1}, 1, CUDA.Mem.DeviceBuffer})
    @ Zygote ~/.julia/packages/Zygote/FPUm3/src/compiler/interface.jl:40
 [17] mul!(result::CuArray{Float32, 1, CUDA.Mem.DeviceBuffer}, op::HvpOperator{typeof(f), Float32, Int64}, v::CuArray{Float32, 1, CUDA.Mem.DeviceBuffer})
    @ Main ~/Documents/research/masters-thesis/CubicNewton/experiments/test.jl:34
 [18] top-level scope
    @ ~/Documents/research/masters-thesis/CubicNewton/experiments/test.jl:49
 [19] include(fname::String)
    @ Base.MainInclude ./client.jl:444
 [20] top-level scope
    @ REPL[14]:1
 [21] top-level scope
    @ ~/.julia/packages/CUDA/nYggH/src/initialization.jl:52

```

I know there is quite a lot of information here, so if I can clarify or add anything let me know. Thanks again for the help!

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 27, 2022, 8:21pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/2 "2022-01-27T20:21:24Z")

</div>

you probably need to strictly type the Duals. Did you try the SparseDiffTools implementations?

[https://github.com/JuliaDiff/SparseDiffTools.jl/blob/master/src/differentiation/jaches\_products\_zygote.jl#L27-L42](https://github.com/JuliaDiff/SparseDiffTools.jl/blob/master/src/differentiation/jaches_products_zygote.jl#L27-L42)

---

<div class="post-metadata">

**Author:** ![bad\_at\_math](https://avatars.discourse-cdn.com/v4/letter/b/51bf81/32.png) [@bad\_at\_math](https://discourse.julialang.org/u/bad_at_math)\
**Post date:** [January 27, 2022, 11:05pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/3 "2022-01-27T23:05:44Z")

</div>

@ChrisRackauckas It seems like those SparseDiffTools functions are only exported for a particular, kind of old version of Zygote ([https://github.com/JuliaDiff/SparseDiffTools.jl/blob/master/src/SparseDiffTools.jl#L58](https://github.com/JuliaDiff/SparseDiffTools.jl/blob/master/src/SparseDiffTools.jl#L58)). Is there a reason for this, and/or is it tied to a current SparseDiffTools or Zygote issue?

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 27, 2022, 11:12pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/4 "2022-01-27T23:12:03Z")

</div>

No, they support the most recent version.

---

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 27, 2022, 11:27pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/5 "2022-01-27T23:27:22Z")

</div>

My operator is heavily based on SparseDiffTools, but I haven’t tried the implementation there. When you say strictly typing the Duals do you mean writing `Dual{typeof(ForwardDiff.Tag(DeivVecTag,eltype(x))),eltype(x),1}` as opposed to `Dual{Nothing, T, 1}`? In other words just specifying the type `T`? I am also not quite sure I understand the idea behind the tag, what should I be putting there?

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 28, 2022, 2:11am UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/6 "2022-01-28T02:11:19Z")

</div>

The full type specification can improve inference for Duals.

Though reading this again, the issue is a cudnn one, in that cudnn is not compatible with dual numbers. I think you’d have to use a different form of forward differentiation to get around that, or a different dense kernel.

---

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 28, 2022, 3:44am UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/7 "2022-01-28T03:44:46Z")

</div>

To clarify, ForwardDiff (or Duals in particular) are compatible with CUDA arrays, but not with cuDNN? Without asking too much of you, can you comment on possible alternatives? I think Zygote has some hidden functionality for forward mode AD, and there is the fabled Diffractor.jl?

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 28, 2022, 4:00am UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/8 "2022-01-28T04:00:20Z")

</div>

> [@RS-Coop](#):
>
> I think Zygote has some hidden functionality for forward mode AD, and there is the fabled Diffractor.jl?

You’d need to check for an frule on this op.

---

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 30, 2022, 4:54pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/9 "2022-01-30T16:54:54Z")

</div>

Ok, I think I understand the issue.

When I remove the sigmoid activation from my model the code runs fine on the GPU because there is no issue with an frule being defined for the cuDNN sigmoid. Obviously the hvp will be trivially zero in this case, but it still works.

If I wanted to use a sigmoid then it seems I have two options:

- Define an frule for the cuDNN sigmoid
- Make my own sigmoid function which can then be differentiated using lower level frules even on the GPU

Is that correct?

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 30, 2022, 5:03pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/10 "2022-01-30T17:03:23Z")

</div>

I haven’t dug into the details of your case, but it certainly sounds like that’s the right track. I’m not sure it’s using the fused sigmoid rule, but it could be, in which case that would need an frule.

---

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 30, 2022, 5:10pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/11 "2022-01-30T17:10:08Z")

</div>

Things seem to be working by defining my own activation functions. Thanks for the help here!

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 30, 2022, 6:09pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/12 "2022-01-30T18:09:54Z")

</div>

Open an issue so that future users don’t have to workaround it.

---

<div class="post-metadata">

**Author:** ![RS-Coop](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/rs-coop/32/32175_2.png) [@RS-Coop](https://discourse.julialang.org/u/RS-Coop)\
**Post date:** [January 30, 2022, 10:14pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/13 "2022-01-30T22:14:10Z")

</div>

With ForwardDiff, Flux, or NNlibCUDA? Seems like one of the last two, maybe Flux because they are trying to be AD agnostic.

---

<div class="post-metadata">

**Author:** ![ChrisRackauckas](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/chrisrackauckas/32/77_2.png) [@ChrisRackauckas](https://discourse.julialang.org/u/ChrisRackauckas)\
**Post date:** [January 31, 2022, 2:28pm UTC](https://discourse.julialang.org/t/hessian-vector-products-on-gpu-using-forwarddiff-zygote-and-flux/75306/14 "2022-01-31T14:28:10Z")

</div>

I think it’s an NNlibCUDA issue.
