# How to avoid "unsupported dynamic function invocation" in CUDA with nested gradients

**URL:** <https://discourse.julialang.org/t/how-to-avoid-unsupported-dynamic-function-invocation-in-cuda-with-nested-gradients/127982>\
**Category:** GPU\
**Tags:** question\
**Created:** [April 11, 2025, 3:18pm UTC](https://discourse.julialang.org/t/how-to-avoid-unsupported-dynamic-function-invocation-in-cuda-with-nested-gradients/127982 "2025-04-11T15:18:28Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![andrewrosemberg](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewrosemberg/32/9902_2.png) [@andrewrosemberg](https://discourse.julialang.org/u/andrewrosemberg)\
**Post date:** [April 11, 2025, 3:18pm UTC](https://discourse.julialang.org/t/how-to-avoid-unsupported-dynamic-function-invocation-in-cuda-with-nested-gradients/127982/1 "2025-04-11T15:18:28Z")

</div>

I am trying to train a model in Flux in which the loss has a nested gradient. I know I should avoid dynamic function invocation, but I am unsure how. The following code works on CPU but not on CUDA/GPU:

```julia
using Flux
using Zygote

device = cpu # or gpu
λ = Float32(0.1)

X = Float32.(rand(30, 20)) |> device
y = Float32.(rand(10, 20)) |> device
dx = Float32.(rand(30, 20)) |> device
dy = Float32.(ones(10, 20)) |> device

model = Chain(
    Dense(30, 10, relu),
    Dense(10, 10, relu),
    Dense(10, 10)
) |> device

# 1) If dy = ∂sum(ŷ)/∂ŷ
loss, grad = Zygote.withgradient(model) do model
    ret = Zygote.withgradient(X) do X
        ŷ = model(X)
        return sum(ŷ)
    end
    return Flux.mse(model(X), y) + λ * Flux.mse(dx, ret.grad[1])
end

# 2) General dy
loss_general, grad_general = Zygote.withgradient(model) do model
    ŷ, pb = Zygote.pullback(X) do X
        model(X)
    end
    return Flux.mse(ŷ, y) + λ * Flux.mse(dx, pb(dy)[1])
end

```

It would be great to get 1 and 2 to work on CUDA, but I am happy with just one.

---

<div class="post-metadata">

**Author:** ![gdalle](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/gdalle/32/27854_2.png) [@gdalle](https://discourse.julialang.org/u/gdalle)\
**Post date:** [April 11, 2025, 5:57pm UTC](https://discourse.julialang.org/t/how-to-avoid-unsupported-dynamic-function-invocation-in-cuda-with-nested-gradients/127982/2 "2025-04-11T17:57:41Z")

</div>

Maybe Lux.jl is a bit better at handling nested autodiff naturally?

> **[Nested Automatic Differentiation | Lux.jl Docs](https://lux.csail.mit.edu/stable/manual/nested_autodiff)**
>
> Documentation for LuxDL Repositories

---

<div class="post-metadata">

**Author:** ![andrewrosemberg](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/andrewrosemberg/32/9902_2.png) [@andrewrosemberg](https://discourse.julialang.org/u/andrewrosemberg)\
**Post date:** [April 15, 2025, 2:58pm UTC](https://discourse.julialang.org/t/how-to-avoid-unsupported-dynamic-function-invocation-in-cuda-with-nested-gradients/127982/4 "2025-04-15T14:58:32Z")

</div>

Thank you. I followed the instructions there, and it worked!
