# Checkpointing and adjoints for custom kernels

**URL:** <https://discourse.julialang.org/t/checkpointing-and-adjoints-for-custom-kernels/108281>\
**Category:** Modelling & Simulations\
**Tags:** question\
**Created:** [January 3, 2024, 3:48am UTC](https://discourse.julialang.org/t/checkpointing-and-adjoints-for-custom-kernels/108281 "2024-01-03T03:48:38Z")\
**Posts on this page:** 3\
**Page:** 1

<div class="post-metadata">

**Author:** ![smartalecH](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/smartalech/32/32379_2.png) [@smartalecH](https://discourse.julialang.org/u/smartalecH)\
**Post date:** [January 3, 2024, 3:48am UTC](https://discourse.julialang.org/t/checkpointing-and-adjoints-for-custom-kernels/108281/1 "2024-01-03T03:48:39Z")

</div>

I’ve written a PDE finite-difference solver that uses KernelAbstractions.jl to update my stencil each time step.

In my previous life, when I’d write this stuff in C/C++, I’d roll my own adjoints, and then hook them up to an AD package. And if I really needed to, I’d implement some naive checkpointing to alleviate any memory restrictions.

I’ve looked through the various SciML packages, and it seems like there’s a lot of machinery in place to automate the adjoint sensitivity and even checkpointing. But it also seems like I have to be “all in” and use the full SciML ecosystem to discretize the PDE etc.

Am I wrong? Is there a way to use my existing kernel code, and wrap it with some of the SciML packages just to handle the adjoint/backpropagation and the checkpointing? What’s the canonical way people typically address this?

To give some context, my stencil code modifies my data arrays in place, so relying on something using Zygote probably wouldn’t work (although I hope I’m wrong).

Thanks!

---

<div class="post-metadata">

**Author:** ![luraess](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/luraess/32/16189_2.png) [@luraess](https://discourse.julialang.org/u/luraess)\
**Post date:** [January 3, 2024, 9:37pm UTC](https://discourse.julialang.org/t/checkpointing-and-adjoints-for-custom-kernels/108281/2 "2024-01-03T21:37:06Z")

</div>

A potential alternative could be to use [Enzyme.jl](https://github.com/EnzymeAD/Enzyme.jl) which has good integration in KernelAbstractions.jl, together with [Checkpointing.jl](https://github.com/Argonne-National-Laboratory/Checkpointing.jl). Setting up that machinery may be more “manual” maybe as such, but could come with some good performance gain and should be GPU compatible as well.

---

<div class="post-metadata">

**Author:** ![smartalecH](https://sea2.discourse-cdn.com/julialang/user_avatar/discourse.julialang.org/smartalech/32/32379_2.png) [@smartalecH](https://discourse.julialang.org/u/smartalecH)\
**Post date:** [January 6, 2024, 12:32am UTC](https://discourse.julialang.org/t/checkpointing-and-adjoints-for-custom-kernels/108281/3 "2024-01-06T00:32:52Z")

</div>

Great suggestion, thanks!
