cfei1994/mosquito-app
0
1abstract type ModelInferencePipeline end2mutable struct MosquitoInference <: ModelInferencePipeline3 pos_track::Vector{Matrix{Float64}} # x_i4 vel_track::Vector{Matrix{Float64}} # v_i5 tobs::Vector{Array{Float64}} # t_i6 model::MosquitoModel # model with functional term7 Θ::Vector{Array{Float64,3}} # F(x, v)8 dVdt::Vector{Matrix{Float64}} # dVdt - F_learned 9 coef_full::Vector{Float64} # w_μ10 bigind::BitVector11 G::Matrix{Float64}12 b::Vector{Float64} # dV/dt - F_learned13 cacheG::Matrix{Float64}14 cacheb::Vector{Float64}15 16 MosquitoInference(robs::Vector{Matrix{Float64}}, vobs::Vector{Matrix{Float64}}, 17 tobs::Vector{Array{Float64}}, model::MosquitoModel) = new(18 robs, vobs, tobs, model,19 Vector{Array{Float64,3}}(undef, length(robs)),20 Vector{Array{Float64,3}}(undef, length(robs)),21 [0.0], trues(1), zeros(1,1), zeros(1), zeros(1,1), zeros(1)22 )23end24 25 26 27# ============= FUNCTIONS for the inference pipeline ==================28 29# ------------- Initialization -------------30function init(P::MosquitoInference)31 32 # get keys 33 allkeys, learnkey = internal_dim_check(P)34 nterm = sum( [length(P.model.funcs[learnkey_]) for learnkey_ in learnkey] )35 36 # init Θ (using the first slice to store the velocity trajectories)37 # init dVdt38 for i in eachindex(P.Θ)39 P.Θ[i] = zeros(size(P.vel_track[i])...,nterm)40 # dts = P.tobs[i][2]-P.tobs[i][1]41 dts = @view(P.tobs[i][2:end,:]) .- @view(P.tobs[i][1:end-1,:]) 42 P.dVdt[i] = (@view(P.vel_track[i][2:end,:]) .- @view(P.vel_track[i][1:end-1,:]) ) ./ dts43 end44 45 # update dVdt46 learnedkeys = allkeys[[P.model.islearned[key] for key in allkeys]]47 for lkey in learnedkeys48 println("learned force: $(lkey)")49 _update_dVdt!(P.dVdt, P.pos_track, P.vel_track, P.model.funcs[lkey], P.model.params[lkey], P.model.coeffs[lkey], 50 lkey, P.model.ndim, P.model.bfields[lkey])51 end52 53 # init coeffs54 P.coef_full = zeros(nterm)55 56 # init bigind57 P.bigind = trues(nterm)58 59 # init b60 _init_b!(P)61 62 # init G63 P.G = zeros(length(P.b), nterm)64 65 return nothing66end67 68 69 70function internal_dim_check(P::MosquitoInference)71 72 # check the input data has the same dimension has the model73 any(size.(P.pos_track,2) .!= P.model.ndim) && (error("DimensionError: the dimension of the model must be the same as that of the positional data"));74 any(size.(P.vel_track,2) .!= P.model.ndim) && (error("DimensionError: the dimension of the model must be the same as that of the velocity data"));75 (length(P.pos_track) != length(P.vel_track)) && (error("DimensionError: pos and vel data must have the same number of particles"))76 77 78 # check only one force term in in the learning process79 keys = [fieldnames(typeof(P.model.islearning))...]80 islearning_vec = [P.model.islearning[k] for k in keys]81 (sum(islearning_vec) != 1) && (error("need at least one and only one force term to be inferred")) 82 83 return keys, keys[islearning_vec]84 85end86 87 88 89 90function _update_dVdt!(dVdt::Vector{Matrix{Float64}}, pos::Vector{Matrix{Float64}}, vel::Vector{Matrix{Float64}},91 funcs, params::Vector{Float64}, coeffs::Vector{Float64}, key::Symbol, ndim::Int, get_bfield!::Function)92 93 any(size.(dVdt,1) .!= (size.(pos,1).-1) ) && (error("DimensionError: size(dVdt,1) == size(pos,1)-1 not satisfied"))94 95 ri_vec = zeros(ndim)96 vi_vec = zeros(ndim)97 u_vec = zeros(ndim)98 b_vec = zeros(ndim)99 100 for pid in eachindex(pos)101 102 for t in 1:size(dVdt[pid],1)103 104 for k in 1:ndim105 ri_vec[k] = pos[pid][t,k]106 vi_vec[k] = vel[pid][t,k]107 end108 109 # a_mag110 vi_mag = mynorm(vi_vec)111 vi_vec ./= vi_mag112 113 # b_mag114 b_mag = get_bfield!(b_vec, ri_vec);115 ab_dot = 0.0;116 for k in 1:ndim117 ab_dot += vi_vec[k] * b_vec[k];118 end119# ab_dot = dot(vi_vec, b_vec);120 121 fmag = 0.0 122 for idx in eachindex(funcs)123 (typeof(funcs[idx])<:Ψ1) && (fmag = term_mag(funcs[idx], vi_mag, params[1]))124 (typeof(funcs[idx])<:Ψ2) && ( fmag = term_mag(funcs[idx], vi_mag, b_mag, ab_dot, params[1], params[2]) ) 125 (funcs[idx].uvec == :ahat) && (copy!(u_vec, vi_vec))126 (funcs[idx].uvec == :a) && (copy!(u_vec, vi_vec); u_vec .*= vi_mag;)127 (funcs[idx].uvec == :bhat) && (copy!(u_vec, b_vec))128 (funcs[idx].uvec == :b) && (copy!(u_vec, b_vec); u_vec .*= b_mag;)129 (funcs[idx].uvec == :bhatorth) && ( u_vec .= (b_vec .- vi_vec .* ab_dot) ./ vi_mag )130 (funcs[idx].uvec == :bhatorth2) && ( u_vec .= (b_vec .- vi_vec .* ab_dot) )131 for k in 1:ndim132 dVdt[pid][t,k] -= coeffs[idx] * fmag * u_vec[k]133 end134 end135 136 end137 138 end139 140 return nothing141end142 143 144function _init_b!(P::MosquitoInference)145 row_array = length.(P.dVdt)146 P.b = zeros(sum(row_array))147 count = 0 148 for ip in eachindex(P.dVdt)149 copyto!(@view(P.b[count+1:count+row_array[ip]]),150 P.dVdt[ip])151 count += row_array[ip]152 end153 return nothing154end155 156 157 158function build_theta!(Theta_s::Vector{Array{Float64,3}}, pos::Vector{Matrix{Float64}}, vel::Vector{Matrix{Float64}},159 funcs, params::Vector{Float64}, key::Symbol, ndim::Int, get_bfield!::Function)160 161 ri_vec = zeros(ndim)162 vi_vec = zeros(ndim)163 u_vec = zeros(ndim)164 b_vec = zeros(ndim)165 166 for pid in eachindex(pos)167 168 fill!(Theta_s[pid], 0.0)169 170 for t in 1:size(pos[pid],1)171 172 for k in 1:ndim173 ri_vec[k] = pos[pid][t,k]174 vi_vec[k] = vel[pid][t,k]175 end176 177 # a_mag178 vi_mag = mynorm(vi_vec)179 vi_vec ./= vi_mag180 181 # b_mag182 b_mag = get_bfield!(b_vec, ri_vec);183 ab_dot = 0.0;184 for k in 1:ndim185 ab_dot += vi_vec[k] * b_vec[k];186 end187# ab_dot = dot(vi_vec, b_vec);188 189 fmag = 0.0 190 for idx in eachindex(funcs)191 # fmag192 (typeof(funcs[idx])<:Ψ1) && ( fmag = term_mag(funcs[idx], vi_mag, params[1]))193 (typeof(funcs[idx])<:Ψ2) && ( fmag = term_mag(funcs[idx], vi_mag, b_mag, ab_dot, params[1], params[2]) ) 194 195 # uvec196 (funcs[idx].uvec == :ahat) && (copy!(u_vec, vi_vec))197 (funcs[idx].uvec == :a) && (copy!(u_vec, vi_vec); u_vec .*= vi_mag;)198 (funcs[idx].uvec == :bhat) && (copy!(u_vec, b_vec))199 (funcs[idx].uvec == :b) && (copy!(u_vec, b_vec); u_vec .*= b_mag;)200 (funcs[idx].uvec == :bhatorth) && ( u_vec .= (b_vec .- vi_vec .* ab_dot) ./ vi_mag )201 (funcs[idx].uvec == :bhatorth2) && ( u_vec .= (b_vec .- vi_vec .* ab_dot) )202 for k in 1:ndim203 Theta_s[pid][t, k, idx] += fmag * u_vec[k]204 end205 end206 207 end208 209 end210end211 212 213@inline _update_theta!(P::MosquitoInference, model::MosquitoModel, sym::Symbol) = build_theta!(P.Θ, P.pos_track, 214 P.vel_track, model.funcs[sym], model.params[sym], sym, model.ndim, model.bfields[sym])215 216 217 218function _update_G!(G::Matrix{Float64}, Θ::Vector{Array{Float64,3}})219 220 count=0221 222 for ip in eachindex(Θ)223 224 L1 = size(Θ[ip],1) - 1225 L2 = size(Θ[ip],2)226 nrows= L1*L2227 228 for n in axes(G,2)229 copyto!( @view( G[count+1:count+nrows,n] ),230 @view( Θ[ip][1:end-1,:,n])231 )232 end233 234 count += nrows235 236 end237 238end239 240 241 242function sparse_bayesian_fit(self::MosquitoInference, params_new, sym::Symbol; opt_args...)243 244 copy!(self.model.params[sym], params_new)245 _update_theta!(self, self.model, sym)246 _update_G!(self.G, self.Θ)247 sbl_res = SBL(self.G, self.b; opt_args...)248 249 return sbl_res250 251end252 253 254 255function SBL(Phi::Matrix{Float64}, Y::Vector{Float64}; 256 MAX_ITERS=100, EPSILON=1e-6, lambda=1.0, gamma=0.5*ones(size(Phi,2)))257 258 # fitting Y = Phi * mu + GaussianNoise(lambda)259 260 N, M = size(Phi)261 @assert N == length(Y)262 263 mu = zeros(M)264 mu_old = zeros(M)265 266 PhiT_Phi = Phi'*Phi;267 G_inv = Diagonal(zeros(M));268 Sigma = zeros(M,M)269 Xi = zeros(M, N);270 deltaY = similar(Y)271 272 count = 0273 274 while true275 276 copyto!(mu_old, mu)277 278 G_inv.diag .= 1 ./ gamma279 280 Sigma .= PhiT_Phi ./ lambda .+ G_inv .+ 1e-8281 282 Q = lu!(Sigma);283 Sigma .= Q \ I;284 285# LinearAlgebra.inv!(cholesky!(Sigma))286# Sigma .= inv(Sigma)287 288 289 BLAS.gemm!('N','T',1/lambda, Sigma, Phi, false, Xi)290 mul!( mu, Xi, Y)291 292 copyto!(deltaY, Y)293 BLAS.gemm!('N','N',-1.0, Phi, mu, true, deltaY)294 295 for k in eachindex(gamma)296 gamma[k] = mu[k]*mu[k] + Sigma[k,k]297 end298 299 lambda = sum(abs2, deltaY)300 # lambda /= (N-sum(gamma))301 lambda /= N302 303 count += 1304 (count >= MAX_ITERS) && (break)305 (all(abs.(mu_old .- mu) .< EPSILON)) && (break)306 307 end308 309 return 1/2*log(lambda), mu310 311end