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:1637Let 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(θ)
nothingf(θ) = 314First, 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)
nothingderivative_estimate(f, θ, StochasticAD.ForwardAlgorithm(PrunedFIsBackend()); direction = u) = -4.0Now, 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.0Let'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)
nothingForward mode: -1205.9251 ± 12.082082198756396
Reverse mode: -1198.9259 ± 12.043949753543117This page was generated using Literate.jl.