# Strategies to use ReverseDiff.jl with NamedTuples (or ComponentArrays)

**URL:** <https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760>\
**Category:** General Usage\
**Tags:** question\
**Created:** [March 30, 2022, 6:08pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760 "2022-03-30T18:08:03Z")\
**Posts on this page:** 6\
**Page:** 1

<div class="post-metadata">

**Author:** ![roualdes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roualdes/32/20716_2.png) [@roualdes](https://discourse.julialang.org/u/roualdes)\
**Post date:** [March 30, 2022, 6:08pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/1 "2022-03-30T18:08:03Z")

</div>

It seems like ReverseDiff.jl can be used on functions that accept NamedTuples, for instance in [LogDensityProblems.jl](https://github.com/tpapp/LogDensityProblems.jl) and [TransformVariables.jl](https://github.com/tpapp/TransformVariables.jl), but I can’t figure out the details. Will you help me understand how I could change my attempt below to use ReverseDiff.jl with say the function `g`, instead of just `f`?

Here’s an example showing that LogDensityProblems and TransformVariables work well with ReverseDiff, and also highlights where I’m stuck in my attempt.

```julia
using TransformVariables
using LogDensityProblems
using ReverseDiff

g(θ) = -0.5 * θ.x' * θ.x
f(x) = -0.5 * x' * x

el = TransformedLogDensity(as((x = as(Array, 2),)), g);
adg = ADgradient(:ReverseDiff, el);

# my attempt
struct RAD{F, T, R}
    lp::F
    tape::T
    result::R
end

function RAD(lp, x)
    tape = ReverseDiff.compile(ReverseDiff.GradientTape(lp, (x,)))
    res = map(ReverseDiff.DiffResults.GradientResult, (similar(x), ))
    return RAD(lp, tape, res)
end

x = randn(2);
adf = RAD(f, similar(x));

LogDensityProblems.logdensity_and_gradient(adg, x)
ReverseDiff.gradient!(adf.result, adf.tape, (x,))

RAD(g, (x = x,)) # errors

```

The immediate error is that there is “no method matching similar(::NamedTuple{(:x,), Tuple{Vector{Float64}}})”, but I believe this is just the first error of many to follow.

Would you help me understand the strategy used in LogDensityProblems.jl and TransformVariables.jl? I tried reading the code, but was stumped by the function [https://github.com/tpapp/TransformVariables.jl/blob/e6efa6ac266a3bf5d5fd3b26be443bb35391f1c9/src/aggregation.jl#L57](https://github.com/tpapp/TransformVariables.jl/blob/e6efa6ac266a3bf5d5fd3b26be443bb35391f1c9/src/aggregation.jl#L57)

Thanks in advance.

---

<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:** [March 30, 2022, 6:33pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/2 "2022-03-30T18:33:37Z")

</div>

ReverseDiff is not friends with many structs. I would use a different AD if you need to handle such cases.

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [March 30, 2022, 7:23pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/3 "2022-03-30T19:23:09Z")

</div>

> [@ChrisRackauckas](#):
>
> ReverseDiff is not friends with many structs.

I plan to make the introduction, one of those weekends.

---

<div class="post-metadata">

**Author:** ![mohamed82008](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/mohamed82008/32/18171_2.png) [@mohamed82008](https://discourse.julialang.org/u/mohamed82008)\
**Post date:** [March 30, 2022, 7:26pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/4 "2022-03-30T19:26:55Z")

</div>

The standard approach is to flatten the struct and take the gradient wrt to a vector. Then unflatten the gradient as a post-processing step.

---

<div class="post-metadata">

**Author:** ![jonniedie](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/jonniedie/32/12842_2.png) [@jonniedie](https://discourse.julialang.org/u/jonniedie)\
**Post date:** [March 31, 2022, 1:08am UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/5 "2022-03-31T01:08:17Z")

</div>

Is this what you’re looking for?

```julia
using ComponentArrays, ReverseDiff

g(θ) = -0.5 * θ.x' * θ.x
θ = ComponentArray(x=randn(2))

ReverseDiff.gradient(g, θ)
# ComponentVector{Float64}(x = [0.16737302092295353, -0.4702876184629478])

```

edit: I’m not really sure what the `TransformedLogDensity` stuff is doing (I’m not really familiar with TransformVariables.jl), but it seems that `logdensity_and_gradient` is just calculating the value and gradient of `g(θ)`. If that’s the case, you don’t really need anything besides `ReverseDiff.gradient`, I think.

---

<div class="post-metadata">

**Author:** ![roualdes](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/roualdes/32/20716_2.png) [@roualdes](https://discourse.julialang.org/u/roualdes)\
**Post date:** [March 31, 2022, 8:21pm UTC](https://discourse.julialang.org/t/strategies-to-use-reversediff-jl-with-namedtuples-or-componentarrays/78760/6 "2022-03-31T20:21:36Z")

</div>

@jonniedie indeed, ComponentArrays works within my attempts to pre-compile the tape too. My mistake. ~~I’ll remove my parenthetical remark from the title of this thread.~~ edit: Turns out I won’t edit the title, cause either I can’t or I don’t know how.

Thanks all for your help. Looking forward to seeing the outcome of ReverseDiff introduced to NamedTuples 🙂
