Skip to content

Instantly share code, notes, and snippets.

@mschauer
Created May 4, 2021 15:36
Show Gist options
  • Select an option

  • Save mschauer/da7bddafa958e674f666df3201c3375e to your computer and use it in GitHub Desktop.

Select an option

Save mschauer/da7bddafa958e674f666df3201c3375e to your computer and use it in GitHub Desktop.
Zero cost abstraction Maruyama solver
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, ())
@mschauer

mschauer commented May 4, 2021

Copy link
Copy Markdown
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

@fmeulen

fmeulen commented May 4, 2021 via email

Copy link
Copy Markdown

@mschauer

mschauer commented May 4, 2021

Copy link
Copy Markdown
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