# Help with Jacobian vector product to get natural gradient

**URL:** <https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115>\
**Category:** Probabilistic Programming\
**Tags:** forwarddiff, natural-gradient\
**Created:** [December 2, 2020, 2:03pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115 "2020-12-02T14:03:03Z")\
**Posts on this page:** 17\
**Page:** 1

<div class="post-metadata">

**Author:** ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)\
**Post date:** [December 2, 2020, 2:03pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/1 "2020-12-02T14:03:03Z")

</div>

Hi I am trying to reproduce an algorithm based on Natural gradient computations ([Natural Gradients in Practice - Salimbeni et al 2018](http://arxiv.org/abs/1803.09151)).  
The key computation is that you can get rid of the inverse Fisher Information matrix by replacing it with transformation derivatives, here is the relevant extract ;

 ![2020-12-02_14-59](https://global.discourse-cdn.com/julialang/original/3X/d/7/d770334af74914f7e07ddf0f3e66cccac84b91b3.png)

The whole talk about reverse-mode relevant, as forward mode is available in Julia. However I was wondering if there was a way to compute this quantity in one pass instead of computing first \frac{\partial \xi}{\partial \theta} and then \frac{\partial \mathcal{L}}{\partial \eta}.

For precision here \xi are an arbitrary representation of the variational parameters, \theta are the natural parameters and \eta are the expectation parameters (first and second moment)

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 2:54pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/2 "2020-12-02T14:54:24Z")

</div>

👋

> However I was wondering if there was a way to compute this quantity in one pass instead of computing first ∂ξ∂θ and then ∂L∂η .

I’m a little confused by this statement, since at no point do you ever actually instantiate ∂ξ∂θ. So could you elaborate a little on what you mean by this?

---

<div class="post-metadata">

**Author:** ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)\
**Post date:** [December 2, 2020, 2:58pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/3 "2020-12-02T14:58:06Z")

</div>

👋

Well if we take a concrete example where \xi = (\mu, L) where q = \mathcal{N}(\mu, LL^\top), you still need to compute \frac{\partial (\mu, L)}{\partial \theta} right ?

---

<div class="post-metadata">

**Author:** ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)\
**Post date:** [December 2, 2020, 3:02pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/4 "2020-12-02T15:02:22Z")

</div>

I actually wrote a basic script for it to make more clear what I am doing

```julia
    using Flux: destructure
    using LinearAlgebra
    using ForwardDiff: gradient, jacobian
    E, to_expec = destructure(meanvar_to_expec(μ, L))
    dL_dexpec = gradient(E) do E
        μ, L = expec_to_meanvar(to_expec(E)...)
        θ = L * randn(length(μ), nSamples) .+ μ
        sum(logπ, eachcol(θ)) / nSamples + logdet(L)
    end

    η, to_nat = destructure(meanvar_to_nat(μ, L))
    dξ_dη = jacobian(η) do η
        vec(nat_to_meanvar(to_nat(η)...))
    end

    ξ, to_meanvar = destructure((μ, L))
    nat_grad = dξ_dη * dL_dexpec
    Δμ, ΔL = to_meanvar(nat_grad)

```

I used `destructure` out of laziness

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 3:18pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/5 "2020-12-02T15:18:31Z")

</div>

Right, so the point here is that you don’t ever compute the `jacobian(η)` explicitly – instead, to compute the natural gradient w.r.t. the parameters of your preferred parametrisation `ξ` you do something like

```julia
foward_mode_AD(natural_to_meanvar, theta, dL_dexpec)

```

which is equivalent to the jacobian-vector product in your code.

Does this clear things up, or am I missing the point?

---

<div class="post-metadata">

**Author:** ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)\
**Post date:** [December 2, 2020, 3:25pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/6 "2020-12-02T15:25:00Z")

</div>

Ah i think I get it better, you don’t compute the jacobian explicitly, instead you pass `dL_dexpec` as your final value and compute the result directly from there (as a jacobian vector product).  
I cannot find such a function in ForwardDiff.jl, do you know what to use?

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 3:27pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/7 "2020-12-02T15:27:16Z")

</div>

Exactly.

> I cannot find such a function in ForwardDiff.jl, do you know what to use?

I’m actually not entirely sure either – I agree that I can’t immediately see it in the public API. Probably best to ask about ForwardDiff + jvps in #autodiff on Slack for a quick answer.

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 3:39pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/8 "2020-12-02T15:39:13Z")

</div>

Note that if that doesn’t work you can presumably just use the reverse-mode trick from the paper.

---

<div class="post-metadata">

**Author:** ![antoine-levitt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/antoine-levitt/32/4008_2.png) [@antoine-levitt](https://discourse.julialang.org/u/antoine-levitt)\
**Post date:** [December 2, 2020, 3:48pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/9 "2020-12-02T15:48:06Z")

</div>

> Forward-mode automatic differentiation libraries are perhaps less common than reverse-mode, but fortunately there is an elegant way to achieve forward-mode automatic differentiation using reverse-mode differentiation twice

Seriously?! Isn’t forward-mode AD orders of magnitude simpler to implement than reverse-mode?!

For jvp f’(x) \* h, can’t you just differentiate f(x+t h) wrt t?

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 4:09pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/10 "2020-12-02T16:09:14Z")

</div>

Yeah, it’s more of a comment on AD tools that the ML community tend to use, rather than AD generally – JAX is probably the first major bit of ML tooling that has forwards mode.

But yes, I agree that it’s a bit strange.

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 4:11pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/11 "2020-12-02T16:11:19Z")

</div>

From @oxinabox on slack:

> SHould be easy enough if just construct the dual numbers youself right?
> 
> `duals = map(Dual, x, dx)`

So the idea would be to write something like

```julia
out_duals = natural_to_meanvar(map(Dual, theta, dL_dexpec))

```

and then the desired natural gradient should be contained within the dual bits of `out_duals`.

---

<div class="post-metadata">

**Author:** ![YingboMa](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/yingboma/32/2181_2.png) [@YingboMa](https://discourse.julialang.org/u/YingboMa)\
**Post date:** [December 2, 2020, 4:36pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/12 "2020-12-02T16:36:21Z")

</div>

The derivative of `f(t * a + c)` wrt `t` at 0 is the Jacobian vector product `J(f, c) * a`.

```julia
julia> using ForwardDiff

julia> foo(x) = [x[1], x[1]*x[3], x[2]^2]
foo (generic function with 1 method)

julia> ForwardDiff.derivative(t->foo([1,2,3] * t + [3, 4, 5]), 0)
3-element Vector{Int64}:
  1
 14
 16

julia> ForwardDiff.jacobian(foo, [3,4,5]) * [1,2,3]
3-element Vector{Int64}:
  1
 14
 16

```

---

<div class="post-metadata">

**Author:** ![willtebbutt](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/willtebbutt/32/6790_2.png) [@willtebbutt](https://discourse.julialang.org/u/willtebbutt)\
**Post date:** [December 2, 2020, 4:37pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/13 "2020-12-02T16:37:42Z")

</div>

> For jvp f’(x) \* h, can’t you just differentiate f(x+t h) wrt t?

Sorry, didn’t appreciated this properly when you first wrote it!

---

<div class="post-metadata">

**Author:** ![theogf](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/theogf/32/1987_2.png) [@theogf](https://discourse.julialang.org/u/theogf)\
**Post date:** [December 2, 2020, 4:39pm UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/14 "2020-12-02T16:39:16Z")

</div>

@willtebbutt to answer your concern with the problems of such a method, here is the kind of optimiser you need to use to make sure everything is okay :

```julia
struct IncreasingRate
    α::Float64 # Maximum learning rate
    γ::Float64 # Convergence rate to the maximum
    state
end

IncreasingRate(α=1.0, γ=1e-8) = IncreasingRate(α, γ, IdDict())

function Optimise.apply!(opt::IncreasingRate, x, g)
    t = get!(()->0, opt.state, x)
    opt.state[x] += 1
    return g .* opt.α * (1 - exp(-opt.γ * t))
end

```

---

<div class="post-metadata">

**Author:** ![MilkshakeForReal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/milkshakeforreal/32/32861_2.png) [@MilkshakeForReal](https://discourse.julialang.org/u/MilkshakeForReal)\
**Post date:** [February 16, 2022, 1:20am UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/15 "2022-02-16T01:20:09Z")

</div>

Is it only computing the product? Or under the hood it computes the Jacobian first, and then computes the product

---

<div class="post-metadata">

**Author:** ![stevengj](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/stevengj/32/71_2.png) [@stevengj](https://discourse.julialang.org/u/stevengj)\
**Post date:** [February 16, 2022, 2:06am UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/16 "2022-02-16T02:06:52Z")

</div>

Forward mode differentiation of `f(x + t h)` with respect to `t` only computes the product (the directional derivative). It does _not_ compute the Jacobian matrix `f'(x)` first and then multiply it by `h` (which would be vastly less efficient).

---

<div class="post-metadata">

**Author:** ![MilkshakeForReal](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/milkshakeforreal/32/32861_2.png) [@MilkshakeForReal](https://discourse.julialang.org/u/MilkshakeForReal)\
**Post date:** [February 18, 2022, 12:15am UTC](https://discourse.julialang.org/t/help-with-jacobian-vector-product-to-get-natural-gradient/51115/17 "2022-02-18T00:15:53Z")

</div>

Thanks for the explaination!
