CoolFace
Apppublic

cfei1994/mosquito-app

sourceHugging Faceupdated 7mo agoView on Hugging Face
0likes
optimization.jl311 linesDownload Raw Back to mosquito_code
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