Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ DensityInterface = "b429d917-457f-4dbc-8f4c-0cc954292b1d"
DocStringExtensions = "ffbed154-4ef7-542d-bbb7-c09d3a79fcae"
FillArrays = "1a297f60-69ca-5386-bcde-b61e274b549b"
LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
OhMyThreads = "67456a42-1dca-4109-a031-0a68de7e3ad5"
ProgressLogging = "33c8b6b6-d38a-422a-b730-caa89a2f386c"
Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c"
SparseArrays = "2f01184e-e22b-5df5-ae63-d93ebab69eaf"
Expand All @@ -30,6 +31,7 @@ Distributions = "0.25"
DocStringExtensions = "0.9"
FillArrays = "1"
LinearAlgebra = "1"
OhMyThreads = "0.8"
ProgressLogging = "0.1"
Random = "1"
SparseArrays = "1"
Expand Down
3 changes: 2 additions & 1 deletion src/HiddenMarkovModels.jl
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,12 @@ module HiddenMarkovModels

using ArgCheck: @argcheck
using Base: RefValue
using Base.Threads: @threads
using ChainRulesCore: ChainRulesCore, NoTangent, RuleConfig, rrule_via_ad
using DensityInterface: DensityInterface, DensityKind, HasDensity, NoDensity, logdensityof
using DocStringExtensions
using FillArrays: AbstractFill, Fill
using LinearAlgebra: Transpose, axpy!, Diagonal, dot, ldiv!, lmul!, mul!, parent
using OhMyThreads: DynamicScheduler, SerialScheduler, tforeach
using ProgressLogging: @withprogress, @logprogress
using Random: Random, AbstractRNG, default_rng
using SparseArrays: AbstractSparseArray, SparseMatrixCSC, nonzeros, nnz, nzrange, rowvals
Expand All @@ -43,6 +43,7 @@ include("utils/lightcategorical.jl")
include("utils/limits.jl")
include("utils/segments.jl")
include("utils/duration.jl")
include("utils/threading.jl")

include("inference/predict.jl")
include("inference/forward.jl")
Expand Down
10 changes: 2 additions & 8 deletions src/inference/forward.jl
Original file line number Diff line number Diff line change
Expand Up @@ -128,14 +128,8 @@ function forward!(
seq_ends::AbstractVectorOrNTuple{Int},
error_if_not_finite::Bool=true,
)
if seq_ends isa NTuple{1}
for k in eachindex(seq_ends)
_forward!(storage, hmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
else
@threads for k in eachindex(seq_ends)
_forward!(storage, hmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
foreach_sequence(seq_ends) do k
_forward!(storage, hmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
return nothing
end
Expand Down
16 changes: 4 additions & 12 deletions src/inference/forward_backward.jl
Original file line number Diff line number Diff line change
Expand Up @@ -82,18 +82,10 @@ function forward_backward!(
seq_ends::AbstractVectorOrNTuple{Int},
transition_marginals::Bool=true,
)
if seq_ends isa NTuple{1}
for k in eachindex(seq_ends)
_forward_backward!(
storage, hmm, obs_seq, control_seq, seq_ends, k; transition_marginals
)
end
else
@threads for k in eachindex(seq_ends)
_forward_backward!(
storage, hmm, obs_seq, control_seq, seq_ends, k; transition_marginals
)
end
foreach_sequence(seq_ends) do k
_forward_backward!(
storage, hmm, obs_seq, control_seq, seq_ends, k; transition_marginals
)
end
return nothing
end
Expand Down
10 changes: 2 additions & 8 deletions src/inference/forward_hsmm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -363,14 +363,8 @@ function forward!(
seq_ends::AbstractVectorOrNTuple{Int},
error_if_not_finite::Bool=true,
)
if seq_ends isa NTuple{1}
for k in eachindex(seq_ends)
_forward!(storage, hsmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
else
@threads for k in eachindex(seq_ends)
_forward!(storage, hsmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
foreach_sequence(seq_ends) do k
_forward!(storage, hsmm, obs_seq, control_seq, seq_ends, k; error_if_not_finite)
end
return nothing
end
Expand Down
21 changes: 11 additions & 10 deletions src/inference/viterbi.jl
Original file line number Diff line number Diff line change
Expand Up @@ -62,8 +62,15 @@ function _viterbi!(
ϕₜ .+= logBₜ
end

ϕₜ₂ = view(ϕ, :, t2)
q[t2] = argmax(ϕₜ₂)
#= A plain loop rather than `argmax`. Inference mistakes for recursion
when called from within the threaded reduction of `foreach_sequence`. =#
iₘ = 1
for i in axes(ϕ, 1)
if ϕ[i, t2] > ϕ[iₘ, t2]
iₘ = i
end
end
q[t2] = iₘ
logL[k] = ϕ[q[t2], t2]
for t in (t2 - 1):-1:t1
q[t] = ψ[q[t + 1], t + 1]
Expand All @@ -83,14 +90,8 @@ function viterbi!(
control_seq::AbstractVector;
seq_ends::AbstractVectorOrNTuple{Int},
) where {R}
if seq_ends isa NTuple{1}
for k in eachindex(seq_ends)
_viterbi!(storage, hmm, obs_seq, control_seq, seq_ends, k)
end
else
@threads for k in eachindex(seq_ends)
_viterbi!(storage, hmm, obs_seq, control_seq, seq_ends, k)
end
foreach_sequence(seq_ends) do k
_viterbi!(storage, hmm, obs_seq, control_seq, seq_ends, k)
end
return nothing
end
Expand Down
23 changes: 6 additions & 17 deletions src/types/controlled_emission_hmm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -174,23 +174,12 @@ function StatsAPI.fit!(
)
(; γ, ξ) = fb_storage

if seq_ends isa NTuple
for k in eachindex(seq_ends)
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2]
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
end
else
@threads for k in eachindex(seq_ends)
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2]
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
foreach_sequence(seq_ends) do k
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2]
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
end

Expand Down
23 changes: 6 additions & 17 deletions src/types/hmm.jl
Original file line number Diff line number Diff line change
Expand Up @@ -61,23 +61,12 @@ function StatsAPI.fit!(
)
(; γ, ξ) = fb_storage
# Fit states
if seq_ends isa NTuple
for k in eachindex(seq_ends)
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2] # use ξ[t2] as scratch space since it is zero anyway
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
end
else
@threads for k in eachindex(seq_ends)
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2] # use ξ[t2] as scratch space since it is zero anyway
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
foreach_sequence(seq_ends) do k
t1, t2 = seq_limits(seq_ends, k)
scratch = ξ[t2] # use ξ[t2] as scratch space since it is zero anyway
fill!(scratch, zero(eltype(scratch)))
for t in t1:(t2 - 1)
scratch .+= ξ[t]
end
end
fill!(hmm.init, zero(eltype(hmm.init)))
Expand Down
13 changes: 13 additions & 0 deletions src/utils/threading.jl
Original file line number Diff line number Diff line change
@@ -0,0 +1,13 @@
"""
$(SIGNATURES)

Call `f(k)` for every sequence index `k` in `seq_ends`, spreading sequences across tasks.

A single sequence given as an `NTuple{1}` runs serially, which keeps that path type-stable and
allocation-free.
"""
function foreach_sequence(f::F, seq_ends::AbstractVectorOrNTuple{Int}) where {F}
scheduler = seq_ends isa NTuple{1} ? SerialScheduler() : DynamicScheduler()
tforeach(f, eachindex(seq_ends); scheduler)
return nothing
end
8 changes: 7 additions & 1 deletion test/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f"
Documenter = "e30172f5-a6a5-5a46-863b-614d45cd2de4"
Enzyme = "7da242da-08ed-463a-9acd-ee780be4f1d9"
ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210"
HMMTest = "619c5ee3-2be3-4444-b95f-50aebd0fbf42"
HiddenMarkovModels = "84ca31d5-effc-45e0-bfda-5a68cd981f47"
JET = "c3a54625-cd67-489e-a8e7-0a5a0ff4e31b"
JuliaFormatter = "98e50ef6-434e-11e9-1051-2b60c6c9e899"
LinearAlgebra = "37e2e46d-f89d-539d-b4ee-838fcccc9c8e"
Expand All @@ -25,6 +27,10 @@ StatsAPI = "82ae8749-77ed-4fe6-ae5f-f523153014b0"
Test = "8dfed614-e22c-5e08-85e1-65c5234f0b40"
Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f"

[sources]
HMMTest = {path = "../libs/HMMTest"}
HiddenMarkovModels = {path = ".."}

[compat]
Distributions = "0.25.0 - 0.25.128"
JuliaFormatter = "1.0.62"
JuliaFormatter = "1.0.62"
Loading