Simple reverse mode example

Load our packages

using StochasticAD
using Distributions
using Enzyme
using LinearAlgebra
┌ Error: Error during loading of extension OptimisersEnzymeCoreExt of Optimisers, use `Base.retry_load_extensions()` to retry.
│   exception =
│    1-element ExceptionStack:
│    ArgumentError: Package OptimisersEnzymeCoreExt [0dd6cb9a-46b7-5edf-8cb9-8090bdf2dd56] is required but does not seem to be installed:
│     - Run `Pkg.instantiate()` to install all recorded dependencies.
│
│    Stacktrace:
│      [1] __require_prelocked(pkg::Base.PkgId, env::Nothing)
│        @ Base ./loading.jl:2717
│      [2] _require_prelocked(uuidkey::Base.PkgId, env::Nothing)
│        @ Base ./loading.jl:2597
│      [3] _require_prelocked
│        @ ./loading.jl:2591 [inlined]
│      [4] run_extension_callbacks(extid::Base.ExtensionId)
│        @ Base ./loading.jl:1627
│      [5] run_extension_callbacks(pkgid::Base.PkgId)
│        @ Base ./loading.jl:1664
│      [6] run_package_callbacks(modkey::Base.PkgId)
│        @ Base ./loading.jl:1480
│      [7] _require_search_from_serialized(pkg::Base.PkgId, sourcepath::String, build_id::UInt128, stalecheck::Bool; reasons::Dict{String, Int64}, DEPOT_PATH::Vector{String})
│        @ Base ./loading.jl:2206
│      [8] _require_search_from_serialized
│        @ ./loading.jl:2081 [inlined]
│      [9] __require_prelocked(pkg::Base.PkgId, env::String)
│        @ Base ./loading.jl:2729
│     [10] _require_prelocked(uuidkey::Base.PkgId, env::String)
│        @ Base ./loading.jl:2597
│     [11] macro expansion
│        @ ./loading.jl:2525 [inlined]
│     [12] macro expansion
│        @ ./lock.jl:376 [inlined]
│     [13] __require(into::Module, mod::Symbol)
│        @ Base ./loading.jl:2489
│     [14] require
│        @ ./loading.jl:2465 [inlined]
│     [15] eval_import_path
│        @ ./module.jl:36 [inlined]
│     [16] eval_import_path_all(at::Module, path::Expr, keyword::String)
│        @ Base ./module.jl:60
│     [17] _eval_using(to::Module, path::Expr)
│        @ Base ./module.jl:137
│     [18] top-level scope
│        @ reverse_demo.md:22
│     [19] eval(m::Module, e::Any)
│        @ Core ./boot.jl:489
│     [20] #68
│        @ ~/.julia/packages/Documenter/13nbQ/src/expander_pipeline.jl:919 [inlined]
│     [21] cd(f::Documenter.var"#68#69"{Module, Expr}, dir::String)
│        @ Base.Filesystem ./file.jl:112
│     [22] #66
│        @ ~/.julia/packages/Documenter/13nbQ/src/expander_pipeline.jl:918 [inlined]
│     [23] (::IOCapture.var"#12#13"{Type{InterruptException}, Documenter.var"#66#67"{Documenter.Page, Module, Expr}, IOContext{Base.PipeEndpoint}, IOContext{Base.PipeEndpoint}, Base.PipeEndpoint, Base.PipeEndpoint})()
│        @ IOCapture ~/.julia/packages/IOCapture/MR051/src/IOCapture.jl:170
│     [24] with_logstate(f::IOCapture.var"#12#13"{Type{InterruptException}, Documenter.var"#66#67"{Documenter.Page, Module, Expr}, IOContext{Base.PipeEndpoint}, IOContext{Base.PipeEndpoint}, Base.PipeEndpoint, Base.PipeEndpoint}, logstate::Base.CoreLogging.LogState)
│        @ Base.CoreLogging ./logging/logging.jl:542
│     [25] with_logger(f::Function, logger::Base.CoreLogging.ConsoleLogger)
│        @ Base.CoreLogging ./logging/logging.jl:653
│     [26] capture(f::Documenter.var"#66#67"{Documenter.Page, Module, Expr}; rethrow::Type, color::Bool, passthrough::Bool, capture_buffer::IOBuffer, io_context::Vector{Any})
│        @ IOCapture ~/.julia/packages/IOCapture/MR051/src/IOCapture.jl:167
│     [27] kwcall(::@NamedTuple{rethrow::DataType, color::Bool}, ::typeof(IOCapture.capture), f::Function)
│        @ IOCapture ~/.julia/packages/IOCapture/MR051/src/IOCapture.jl:100
│     [28] runner(::Type{Documenter.Expanders.ExampleBlocks}, node::MarkdownAST.Node{Nothing}, page::Documenter.Page, doc::Documenter.Document)
│        @ Documenter ~/.julia/packages/Documenter/13nbQ/src/expander_pipeline.jl:917
│     [29] dispatch(::Type{Documenter.Expanders.ExpanderPipeline}, ::MarkdownAST.Node{Nothing}, ::Vararg{Any})
│        @ Documenter.Selectors ~/.julia/packages/Documenter/13nbQ/src/utilities/Selectors.jl:170
│     [30] expand(doc::Documenter.Document)
│        @ Documenter ~/.julia/packages/Documenter/13nbQ/src/expander_pipeline.jl:60
│     [31] runner(::Type{Documenter.Builder.ExpandTemplates}, doc::Documenter.Document)
│        @ Documenter ~/.julia/packages/Documenter/13nbQ/src/builder_pipeline.jl:224
│     [32] dispatch(::Type{Documenter.Builder.DocumentPipeline}, x::Documenter.Document)
│        @ Documenter.Selectors ~/.julia/packages/Documenter/13nbQ/src/utilities/Selectors.jl:170
│     [33] #101
│        @ ~/.julia/packages/Documenter/13nbQ/src/makedocs.jl:283 [inlined]
│     [34] withenv(::Documenter.var"#101#102"{Documenter.Document}, ::Pair{String, Nothing}, ::Vararg{Pair{String, Nothing}})
│        @ Base ./env.jl:273
│     [35] #99
│        @ ~/.julia/packages/Documenter/13nbQ/src/makedocs.jl:282 [inlined]
│     [36] cd(f::Documenter.var"#99#100"{Documenter.Document}, dir::String)
│        @ Base.Filesystem ./file.jl:112
│     [37] makedocs(; debug::Bool, format::Documenter.HTMLWriter.HTML, kwargs::@Kwargs{sitename::String, authors::String, modules::Vector{Module}, pages::Vector{Pair{String, Any}}, warnonly::Vector{Symbol}})
│        @ Documenter ~/.julia/packages/Documenter/13nbQ/src/makedocs.jl:281
│     [38] kwcall(::@NamedTuple{sitename::String, authors::String, modules::Vector{Module}, format::Documenter.HTMLWriter.HTML, pages::Vector{Pair{String, Any}}, warnonly::Vector{Symbol}}, ::typeof(makedocs))
│        @ Documenter ~/.julia/packages/Documenter/13nbQ/src/makedocs.jl:274
│     [39] top-level scope
│        @ ~/work/StochasticAD.jl/StochasticAD.jl/docs/make.jl:38
│     [40] include(mod::Module, _path::String)
│        @ Base ./Base.jl:309
│     [41] exec_options(opts::Base.JLOptions)
│        @ Base ./client.jl:344
│     [42] _start()
│        @ Base ./client.jl:577
└ @ Base loading.jl:1637

Let us define our target function.

# Define a toy `StochasticAD`-differentiable function for computing an integer value from a string.
string_value(strings, index) = Int(sum(codepoint, strings[index]))
string_value(strings, index::StochasticTriple) = StochasticAD.propagate(index -> string_value(strings, index), index)

function f(θ; derivative_coupling = StochasticAD.InversionMethodDerivativeCoupling())
    strings = ["cat", "dog", "meow", "woofs"]
    index = randst(Categorical(θ); derivative_coupling)
    return string_value(strings, index)
end

θ = [0.1, 0.5, 0.3, 0.1]
@show f(θ)
nothing
f(θ) = 314

First, let's compute the sensitivity of f in a particular direction via forward-mode Stochastic AD.

u = [1.0, 2.0, 4.0, -7.0]
@show derivative_estimate(f, θ, StochasticAD.ForwardAlgorithm(PrunedFIsBackend()); direction = u)
nothing
derivative_estimate(f, θ, StochasticAD.ForwardAlgorithm(PrunedFIsBackend()); direction = u) = -4.0

Now, let's do the same with reverse-mode.

@show derivative_estimate(f, θ, StochasticAD.EnzymeReverseAlgorithm(PrunedFIsBackend(Val(:wins))))
4-element Vector{Float64}:
 -420.0
 -420.0
    0.0
    0.0

Let's verify that our reverse-mode gradient is consistent with our forward-mode directional derivative.

forward() = derivative_estimate(f, θ, StochasticAD.ForwardAlgorithm(PrunedFIsBackend()); direction = u)
reverse() = derivative_estimate(f, θ, StochasticAD.EnzymeReverseAlgorithm(PrunedFIsBackend(Val(:wins))))

N = 40000
directional_derivs_fwd = [forward() for i in 1:N]
derivs_bwd = [reverse() for i in 1:N]
directional_derivs_bwd = [dot(u, δ) for δ in derivs_bwd]
println("Forward mode: $(mean(directional_derivs_fwd)) ± $(std(directional_derivs_fwd) / sqrt(N))")
println("Reverse mode: $(mean(directional_derivs_bwd)) ± $(std(directional_derivs_bwd) / sqrt(N))")
@assert isapprox(mean(directional_derivs_fwd), mean(directional_derivs_bwd), rtol = 3e-2)

nothing
Forward mode: -1205.9251 ± 12.082082198756396
Reverse mode: -1198.9259 ± 12.043949753543117

This page was generated using Literate.jl.