From af934676f03d652dc5836f3da12eaeff231504c5 Mon Sep 17 00:00:00 2001 From: Ryan Senne <50930199+rsenne@users.noreply.github.com> Date: Fri, 25 Sep 2026 17:31:09 -0400 Subject: [PATCH] Prototype OMT --- Project.toml | 2 ++ src/HiddenMarkovModels.jl | 3 ++- src/inference/forward.jl | 10 ++-------- src/inference/forward_backward.jl | 16 ++++------------ src/inference/forward_hsmm.jl | 10 ++-------- src/inference/viterbi.jl | 21 +++++++++++---------- src/types/controlled_emission_hmm.jl | 23 ++++++----------------- src/types/hmm.jl | 23 ++++++----------------- src/utils/threading.jl | 13 +++++++++++++ test/Project.toml | 8 +++++++- 10 files changed, 55 insertions(+), 74 deletions(-) create mode 100644 src/utils/threading.jl diff --git a/Project.toml b/Project.toml index d7bb4c6..b721cd3 100644 --- a/Project.toml +++ b/Project.toml @@ -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" @@ -30,6 +31,7 @@ Distributions = "0.25" DocStringExtensions = "0.9" FillArrays = "1" LinearAlgebra = "1" +OhMyThreads = "0.8" ProgressLogging = "0.1" Random = "1" SparseArrays = "1" diff --git a/src/HiddenMarkovModels.jl b/src/HiddenMarkovModels.jl index 67dd4bb..4a5f5a7 100644 --- a/src/HiddenMarkovModels.jl +++ b/src/HiddenMarkovModels.jl @@ -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 @@ -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") diff --git a/src/inference/forward.jl b/src/inference/forward.jl index 163ed69..301ca2b 100644 --- a/src/inference/forward.jl +++ b/src/inference/forward.jl @@ -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 diff --git a/src/inference/forward_backward.jl b/src/inference/forward_backward.jl index b0b4f72..8db8834 100644 --- a/src/inference/forward_backward.jl +++ b/src/inference/forward_backward.jl @@ -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 diff --git a/src/inference/forward_hsmm.jl b/src/inference/forward_hsmm.jl index 25a8c62..3ac7763 100644 --- a/src/inference/forward_hsmm.jl +++ b/src/inference/forward_hsmm.jl @@ -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 diff --git a/src/inference/viterbi.jl b/src/inference/viterbi.jl index 714727c..dfb32e1 100644 --- a/src/inference/viterbi.jl +++ b/src/inference/viterbi.jl @@ -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] @@ -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 diff --git a/src/types/controlled_emission_hmm.jl b/src/types/controlled_emission_hmm.jl index 64844b1..c135998 100644 --- a/src/types/controlled_emission_hmm.jl +++ b/src/types/controlled_emission_hmm.jl @@ -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 diff --git a/src/types/hmm.jl b/src/types/hmm.jl index 9c1f47b..c6e1775 100644 --- a/src/types/hmm.jl +++ b/src/types/hmm.jl @@ -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))) diff --git a/src/utils/threading.jl b/src/utils/threading.jl new file mode 100644 index 0000000..f582bd4 --- /dev/null +++ b/src/utils/threading.jl @@ -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 diff --git a/test/Project.toml b/test/Project.toml index 183d281..0fd7cd3 100644 --- a/test/Project.toml +++ b/test/Project.toml @@ -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" @@ -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" \ No newline at end of file +JuliaFormatter = "1.0.62"