🚀

Juliaの自動微分パッケージは結局どれが正解なのか?

に公開

Julia Advent Calendar 2025 の1日目です🎄
https://qiita.com/advent-calendar/2025/julia

ここではJulia言語の自動微分パッケージ ForwardDiff.jl, Zygote.jl, Enzyme.jl をそれぞれ直接呼び出す場合と, DifferentiationInterface.jlを経由して呼び出す場合を比較し, どのパッケージを使えばよいのか考えていきます. それぞれの自動微分パッケージの使い方はこちら, DifferentiationInterfaceの使い方はこちらにあるので詳しい説明は省略します. 環境はGoogle Colabです.

インストール

これは1回きりの測定ですが, 今回使用したパッケージのインストールにかかった時間は次の通り. ForwardDiff, Zygote, Enzyme それぞれ 10秒, 110秒, 998秒でした. だいたい10倍ずつ遅い感じですね.

import Pkg
@time Pkg.add("BenchmarkTools")
@time Pkg.add("DifferentiationInterface")
@time Pkg.add("ForwardDiff")
@time Pkg.add("Zygote")
@time Pkg.add("Enzyme")
#  16.608306 seconds (6.67 M allocations: 682.601 MiB, 6.83% gc time)
#  47.924973 seconds (22.00 M allocations: 1.676 GiB, 2.86% gc time, 27.31% compilation time: 41% of which was recompilation)
#  10.175438 seconds (6.56 M allocations: 678.872 MiB, 5.80% gc time)
# 110.377278 seconds (6.75 M allocations: 705.617 MiB, 0.59% gc time, 29 lock conflicts)
# 998.887560 seconds (6.80 M allocations: 694.043 MiB, 0.08% gc time, 0.00% compilation time)

読み込み

これも1回きりの測定ですが, 今回使用するパッケージの読み込みにかかった時間は次の通り. ForwardDiff, Zygote, Enzymeそれぞれ0.15秒, 0.92秒, 41.7秒でした.

@time using BenchmarkTools
@time using DifferentiationInterface
@time import ForwardDiff
@time import Zygote
@time import Enzyme
#  0.158558 seconds (88.36 k allocations: 5.552 MiB)
#  0.015295 seconds (7.12 k allocations: 682.102 KiB)
#  0.152308 seconds (88.76 k allocations: 5.312 MiB, 13.75% compilation time)
#  0.923302 seconds (477.86 k allocations: 30.424 MiB, 5.03% gc time, 11.22% compilation time)
# 41.734993 seconds (43.48 M allocations: 2.355 GiB, 3.67% gc time, 97.67% compilation time: 99% of which was recompilation)

DifferentiationInterfaceなし

以降はBenchmarkTools@btime マクロを用いて計算時間を計測していきます. ForwardDiff, Zygote, Enzymeそれぞれを直接使用して1階微分と勾配を計算しました. 同じパッケージでも1階微分を計算する方法がいくつかあり, Enzymeの場合はForwardとReverseが切り替えられます. 計算時間を計測したところ, 1階微分はForwardDiff.derivativeとZygote.gradientが同率1位でした. 以前に測定した結果では, 勾配ではEnzymeが最速でしたが, 今回はZygoteが最速でした. ただし, 後ほど見るように適切な工夫をするとEnzymeが最速となります.

# 1階微分
@btime ForwardDiff.derivative(x -> x^4, 2.0)
@btime ForwardDiff.derivative(x -> x[1]^4, 2.0)
@btime ForwardDiff.gradient(x -> x[1]^4, [2.0])[1]
@btime Zygote.gradient(x -> x^4, 2.0)[1]
@btime Zygote.gradient(x -> x[1]^4, [2.0])[1][1]
@btime Enzyme.autodiff(Enzyme.Forward, x -> x[1]^4, Enzyme.Duplicated(2.0, 1.0))[1][1]
@btime Enzyme.autodiff(Enzyme.Reverse, x -> x[1]^4, Enzyme.Active(2.0))[1][1]
@btime Enzyme.gradient(Enzyme.Forward, x -> x[1]^4, 2.0)[1]
@btime Enzyme.gradient(Enzyme.Reverse, x -> x[1]^4, 2.0)[1]
  # 2.177 ns (0 allocations: 0 bytes)
  # 1.818 ns (0 allocations: 0 bytes)
  # 567.098 ns (8 allocations: 256 bytes)
  # 1.818 ns (0 allocations: 0 bytes)
  # 17.299 ns (1 allocation: 32 bytes)
  # 9.279 ns (0 allocations: 0 bytes)
  # 9.280 ns (0 allocations: 0 bytes)
  # 9.559 ns (0 allocations: 0 bytes)
  # 9.440 ns (0 allocations: 0 bytes)

# 勾配
@btime ForwardDiff.gradient(x -> x[1] + x[2]^2, [1.0, 1.0])
@btime Zygote.gradient(x -> x[1] + x[2]^2, [1.0, 1.0])[1]
@btime Enzyme.gradient(Enzyme.Forward, x -> x[1] + x[2]^2, [1.0, 1.0])[1]
@btime Enzyme.gradient(Enzyme.Reverse, x -> x[1] + x[2]^2, [1.0, 1.0])[1]
  # 554.941 ns (7 allocations: 320 bytes)
  # 65.047 ns (4 allocations: 160 bytes)
  # 1.195 μs (16 allocations: 576 bytes)
  # 66.888 ns (4 allocations: 160 bytes)

DifferentiationInterfaceあり+工夫なし

ForwardDiff, Zygote, Enzymeそれぞれをバックエンドに使用して1階微分, 2階微分, 勾配, ヘッシアンを計算しました. BenchmarkTools@btime マクロを用いて計算時間を測定したところ, 1階微分はZygote, 2階微分はForwardDiff, 勾配とヘッシアンはZygoteが最速でした. 最速タイムはDifferentiationInterfaceなしの場合とあまり変わらないので, DifferentiationInterfaceを経由しても速度が犠牲になることはほぼないようです.

# 1階微分
@btime derivative(x -> x^4, AutoForwardDiff(), 2.0)
@btime derivative(x -> x^4, AutoZygote(), 2.0)
@btime derivative(x -> x^4, AutoEnzyme(), 2.0)
  # 2.178 ns (0 allocations: 0 bytes)
  # 1.818 ns (0 allocations: 0 bytes)
  # 9.328 ns (0 allocations: 0 bytes)

# 2階微分
@btime second_derivative(x -> x^4, AutoForwardDiff(), 2.0)
@btime second_derivative(x -> x^4, AutoZygote(), 2.0)
@btime second_derivative(x -> x^4, AutoEnzyme(), 2.0)
  # 1.817 ns (0 allocations: 0 bytes)
  # 396.954 μs (1041 allocations: 105.97 KiB)
  # 10.701 ns (0 allocations: 0 bytes)

# 勾配
@btime gradient(x -> x[1] + x[2]^2, AutoForwardDiff(), [1.0, 1.0])
@btime gradient(x -> x[1] + x[2]^2, AutoZygote(), [1.0, 1.0])
@btime gradient(x -> x[1] + x[2]^2, AutoEnzyme(), [1.0, 1.0])
  # 576.115 ns (7 allocations: 320 bytes)
  # 64.898 ns (4 allocations: 160 bytes)
  # 66.968 ns (4 allocations: 160 bytes)

# ヘッシアン
@btime hessian(x -> x[1] + x[2]^2, AutoForwardDiff(), [1.0, 1.0])
@btime hessian(x -> x[1] + x[2]^2, AutoZygote(), [1.0, 1.0])
@btime hessian(x -> x[1] + x[2]^2, AutoEnzyme(), [1.0, 1.0])
  # 724.588 ns (11 allocations: 784 bytes)
  # 418.095 ns (13 allocations: 640 bytes)
  # 3.517 μs (39 allocations: 1.62 KiB)

DifferentiationInterfaceあり+工夫あり

ヘッシアンはあからさまに遅いので, こちらにある方法で高速化した方がよさそうです. ついでに全てのパッケージに対して再計測してみましょう.

  • 通常版
  • キャッシュ版
  • インプレース版
  • キャッシュ+インプレース版

を比較してみます.

for backend in [AutoForwardDiff(), AutoZygote(), AutoEnzyme()]

  println("\n$(backend)\n")

  f(x) = x^4
  x = 2.0

  println("derivative:")
  prep = prepare_derivative(f, backend, x)
  value = [0.0]
  @btime derivative($f, $backend, $x)
  @btime derivative($f, $prep, $backend, $x)
  @btime derivative!($f, $value, $backend, $x)
  @btime derivative!($f, $value, $prep, $backend, $x)

  println("second_derivative:")
  prep = prepare_second_derivative(f, backend, x)
  value = [0.0]
  @btime second_derivative($f, $backend, $x)
  @btime second_derivative($f, $prep, $backend, $x)
  @btime second_derivative!($f, $value, $backend, $x)
  @btime second_derivative!($f, $value, $prep, $backend, $x)

  g(x) = x[1] + x[2]^2
  x = [1.0, 1.0]

  println("gradient:")
  prep = prepare_gradient(g, backend, x)
  value = [0.0, 0.0]
  @btime gradient($g, $backend, $x)
  @btime gradient($g, $prep, $backend, $x)
  @btime gradient!($g, $value, $backend, $x)
  @btime gradient!($g, $value, $prep, $backend, $x)

  println("hessian:")
  prep = prepare_hessian(g, backend, x)
  value = zeros(2,2)
  @btime hessian($g, $backend, $x)
  @btime hessian($g, $prep, $backend, $x)
  @btime hessian!($g, $value, $backend, $x)
  @btime hessian!($g, $value, $prep, $backend, $x)

end

# 
# AutoForwardDiff()
# 
# derivative:
#   5.391 ns (0 allocations: 0 bytes)
#   4.687 ns (0 allocations: 0 bytes)
#   15.057 ns (0 allocations: 0 bytes)
#   15.058 ns (0 allocations: 0 bytes)
# second_derivative:
#   10.418 ns (0 allocations: 0 bytes)
#   10.432 ns (0 allocations: 0 bytes)
#   14.468 ns (0 allocations: 0 bytes)
#   14.416 ns (0 allocations: 0 bytes)
# gradient:
#   479.393 ns (5 allocations: 240 bytes)
#   37.496 ns (2 allocations: 80 bytes)
#   451.132 ns (3 allocations: 160 bytes)
#   13.631 ns (0 allocations: 0 bytes)
# hessian:
#   781.270 ns (9 allocations: 704 bytes)
#   124.630 ns (4 allocations: 224 bytes)
#   766.228 ns (7 allocations: 592 bytes)
#   108.774 ns (2 allocations: 112 bytes)
# 
# AutoZygote()
# 
# derivative:
#   3.960 ns (0 allocations: 0 bytes)
#   3.249 ns (0 allocations: 0 bytes)
#   4.330 ns (0 allocations: 0 bytes)
#   5.697 ns (0 allocations: 0 bytes)
# second_derivative:
#   392.828 μs (1041 allocations: 105.97 KiB)
#   393.426 μs (1041 allocations: 105.97 KiB)
#   397.272 μs (1044 allocations: 106.03 KiB)
#   394.072 μs (1044 allocations: 106.03 KiB)
# gradient:
#   46.598 ns (2 allocations: 80 bytes)
#   47.568 ns (2 allocations: 80 bytes)
#   50.534 ns (2 allocations: 80 bytes)
#   49.812 ns (2 allocations: 80 bytes)
# hessian:
#   394.668 ns (11 allocations: 560 bytes)
#   398.788 ns (11 allocations: 560 bytes)
#   836.000 ns (13 allocations: 592 bytes)
#   833.602 ns (13 allocations: 592 bytes)
# 
# AutoEnzyme()
# 
# derivative:
#   11.286 ns (0 allocations: 0 bytes)
#   9.411 ns (0 allocations: 0 bytes)
#   10.933 ns (0 allocations: 0 bytes)
#   10.944 ns (0 allocations: 0 bytes)
# second_derivative:
#   10.600 ns (0 allocations: 0 bytes)
#   10.575 ns (0 allocations: 0 bytes)
#   17.590 ns (0 allocations: 0 bytes)
#   11.694 ns (0 allocations: 0 bytes)
# gradient:
#   46.990 ns (2 allocations: 80 bytes)
#   46.950 ns (2 allocations: 80 bytes)
#   21.563 ns (0 allocations: 0 bytes)
#   21.894 ns (0 allocations: 0 bytes)
# hessian:
#   3.478 μs (37 allocations: 1.55 KiB)
#   176.659 ns (8 allocations: 352 bytes)
#   3.386 μs (29 allocations: 1.20 KiB)
#   93.340 ns (0 allocations: 0 bytes)

1階微分, 2階微分ではキャッシュやインプレースは逆効果のようですが, 勾配とヘッシアンは大幅に高速化されるようです. ただし, 勾配やヘッシアンであってもZygoteではキャッシュやインプレースは逆効果でした. 今回の最速パッケージは次の通りです.

演算 最速パッケージ 利用方法
1階微分 Zygote キャッシュ利用
2階微分 ForwardDiff 通常利用
勾配 ForwardDiff※ キャッシュ+インプレース版
ヘッシアン Enzyme キャッシュ+インプレース版

※ ここで ForwardDiff が勾配で最速になっている理由は関数形にも依存すると考えられ, ニューラルネットワーク等のネストが深い関数では異なる結果になると考えられます(要検証).

まとめ

以上の結果と過去の記事を踏まえて, 私がパッケージを選定する基準, 理由についてまとめます.

DifferentiationInterfaceの是非

DifferentiationInterfaceを挟んでも速度にはあまり影響しないため, 今後も新たに開発・改良され続けるバックエンドへの移行のコストを踏まえて, DifferentiationInterfaceを利用することとしました. 実際, 以前に測定した結果と比べるとZygoteは大幅に高速化されているようで, 以下のパッケージ選定の判断も覆る可能性があります.

パッケージが1つに限定される場合

DifferentiationInterfaceを挟んで検証した結果から, パッケージが1つに限定される場合は「Enzyme 一択」です. 選定理由については, Enzymeは1階微分と2階微分だと他のパッケージよりやや遅いですが, 勾配とヘッシアンではあまりに速度差が大きいため, パッケージ1つで全て済ませたい場合はEnzyme一択となります. ただし, 1変数関数しか扱わない場合など, 用途が限定される場合はEnzymeにこだわる必要はありません.

パッケージを2つ使ってよい場合

パッケージを2つ使い分けてよい場合に推奨するパッケージ・利用方法は次の通りです.

演算 推奨パッケージ 利用方法
1階微分 ForwardDiff 通常利用
2階微分 ForwardDiff 通常利用
勾配 Enzyme キャッシュ+インプレース版
ヘッシアン Enzyme キャッシュ+インプレース版

選定理由については以下の通りです. Zygoteは1階微分では最速でしたが, 2階微分ではかなり遅くなっていました. 2階微分を使うことを想定すると, ForwardDiffが妥当という結論になります. もちろん, 2階微分を使わないのであればForwardDiffでなくても問題ありません. ForwardDiffは勾配で最速の結果を出していますが, 恐らくマグレで, 今回の関数形とたまたま相性がよかったのだと考えています. 今回の結果と過去の記事の結果を踏まえれば, 勾配やヘッシアンではEnzymeでキャッシュ+インプレース版を利用する方法が最善策であると考えられます.

おわりに

現時点でも既にJuliaの自動微分パッケージは非常に高い水準にありますが, 今後も新たに開発・改良され続けるため, 長期的な視点からDifferentiationInterface.jlを利用することを強く推奨します. 私は ForwardDiff と Enzyme を使い分ける方針にしました. 今回の結論は現時点での妥当かつベストであると考えていますが, 開発状況によってはさらに高速なパッケージが登場しても不思議ではありません.

参考文献

https://juliadiff.org/DifferentiationInterface.jl/DifferentiationInterface/stable/tutorials/basic/#Preparing-for-multiple-gradients

ノートブック

下記のリンクから, Google Colab上で実行できます.

https://colab.research.google.com/drive/1tYlBiqdK00q3r9OGDG8deeZ1nCwo_hI4?usp=sharing

関連記事

https://zenn.dev/ohno/articles/c1aa146fee7d48

https://zenn.dev/ohno/articles/7b4b6a1ec86189

https://zenn.dev/ohno/articles/0d8d24a50316b5

Discussion