Created
May 4, 2021 15:36
-
-
Save mschauer/da7bddafa958e674f666df3201c3375e to your computer and use it in GitHub Desktop.
Zero cost abstraction Maruyama solver
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| using StaticArrays | |
| struct EulerMaruyama! | |
| end | |
| """ | |
| tangent!(du, u, dz, P) | |
| """ | |
| function tangent! | |
| end | |
| function zero_tangent(x, P) # euclidian fallback | |
| zero(x) # check false*similar | |
| end | |
| function zero_tangent(u::Tuple, P) # euclidian fallback | |
| ntuple(length(u)) do i | |
| false*u[i] | |
| end | |
| end | |
| function zero_integrator(Z, P) | |
| tozero!(deepcopy(Z[1]), P) | |
| end | |
| tozero!(x::Union{Number,SArray}, _) = false*x | |
| tozero!(x::Array, _) = Base.fill!(x, false) | |
| function tozero!(u::Tuple, P) | |
| ntuple(length(u)) do i | |
| tozero!(u[i], P) | |
| end | |
| end | |
| function exponential_map!(x::Array, dx, P) | |
| x .+= dx | |
| end | |
| function exponential_map!(x::Union{SArray,Number}, dx, P) | |
| x + dx | |
| end | |
| saveit!(uu, u, P) = push!(uu, deepcopy(u)) | |
| saveit!(::Nothing, u, P) = nothing | |
| endcondition(uu, u, Z, P) = u[1] >= Z[end][1] | |
| endpoint!(u, _) = u | |
| function solve!(solver::EulerMaruyama!, uu, u, Z, P) | |
| du = zero_tangent(u, P) | |
| dz = zero_integrator(Z, P) | |
| uu, u = solve_inner!(solver, uu, u, du, Z, dz, P) | |
| u = endpoint!(u, P) | |
| saveit!(uu, u, P) | |
| uu, u | |
| end | |
| function solve_inner!(solver::EulerMaruyama!, uu, u, du, Z, dz, P) | |
| while true | |
| saveit!(uu, u, P) | |
| dz = dZ!(u, dz, Z, P) | |
| du = tangent!(du, u, dz, P) | |
| u = exponential_map!(u, du, P) | |
| endcondition(uu, u, Z, P) && break | |
| end | |
| uu, u | |
| end | |
| function dZ!(u, dz, Z, P) | |
| i = u[1] | |
| dw = dz[3] | |
| @. dw = Z[i+1][3] - Z[i][3] | |
| (Z[i+1][1] - Z[i][1], Z[i+1][2] - Z[i][2], dw) | |
| end | |
| function tangent!(du, u, dz, P) | |
| @. du[3] = -u[3]*dz[2] + u[3][]*dz[3] | |
| (dz[1], dz[2], du[3]) | |
| end | |
| function exponential_map!(u::Tuple{Int64, Float64, Vector{Float64}}, du::Tuple, P) | |
| x = u[3] | |
| @. x += du[3] | |
| (u[1] + du[1], u[2] + du[2], x) | |
| end | |
| n = 10000 | |
| t = 0:1/n:1.0 | |
| Z = collect(zip(1:n, t, [[w] for w in [0; sqrt(1/n)*cumsum(randn(n))]])) | |
| u = (1, 0.0, [1.0]) | |
| dz = (0, 0.0, [0.0]) | |
| du = (0, 0.0, [0.0]) | |
| solve!(EulerMaruyama!(), nothing, u, Z, ()) | |
| @time solve!(EulerMaruyama!(), nothing, u, Z, ()) | |
| @code_warntype solve!(EulerMaruyama!(), nothing, u, Z, ()) | |
| uu, uT = solve!(EulerMaruyama!(), typeof(u)[], u, Z, ()) | |
Author
1.) Yes, just thought I show how they look
2.) Let me decifer
# u = (i, t, x)
# dz = (1, dt, dW)
# du = (Δi, dt, dX)
function tangent!(du, u, dz, P)
@. du[3] = -u[3]*dz[2] + u[3][]*dz[3] # dX .= -X dt .+ X dW inplace
(dz[1], dz[2], du[3]) # (1, dt, dX)
end
You mean “dz = “dloglik, dt, dW’ I guess.
… On 4 May 2021, at 19:46, Moritz Schauer ***@***.***> wrote:
@mschauer commented on this gist.
1.) Yes, just thought I show how they look
2.) Let me decifer
# u = (i, t, x)
# dz = (1, dt, dW)
# du = (Δi, dt, dX)
function tangent!(du, u, dz, P)
@. du[3] = -u[3]*dz[2] + u[3][]*dz[3] # dX .= -X dt .+ X dW inplace
(dz[1], dz[2], du[3]) # (1, dt, dX)
end
—
You are receiving this because you commented.
Reply to this email directly, view it on GitHub <https://urldefense.proofpoint.com/v2/url?u=https-3A__gist.github.com_da7bddafa958e674f666df3201c3375e-23gistcomment-2D3731359&d=DwMFaQ&c=XYzUhXBD2cD-CornpT4QE19xOJBbRy-TBPLK0X9U2o8&r=7iF-JNHl-g2k0SDAFdfpFNwTGT5g4T8E90C3u4ODfF8&m=l_WQ5UzDCVpJP2VKo38iWywFWyFuiO3YWuCScqXhobI&s=XF3ylBaP-HFnyzftR6Hs_4vui60z6MaOmEqrW9wjNhk&e=>, or unsubscribe <https://urldefense.proofpoint.com/v2/url?u=https-3A__github.com_notifications_unsubscribe-2Dauth_AJTTG32VQV4CAHAOFBRZB7LTMAXIRANCNFSM44DD2LLQ&d=DwMFaQ&c=XYzUhXBD2cD-CornpT4QE19xOJBbRy-TBPLK0X9U2o8&r=7iF-JNHl-g2k0SDAFdfpFNwTGT5g4T8E90C3u4ODfF8&m=l_WQ5UzDCVpJP2VKo38iWywFWyFuiO3YWuCScqXhobI&s=iVauK7zWM3C_LRimuqGFhqmxt5vBwDlAZ-5A_PL2rGU&e=>.
---
Frank van der Meulen
Delft University of Technology
http://dutiosb.twi.tudelft.nl/~meulen/
Author
In my example case, u = (i, t, x) the first entry is an iteration counter so Δi = 1. But you can of course extend to u = (i, t, x, l)` with also a log-likelihood.
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
tangent!, right?