From eb24f130a0a1d4c32cc54e8d77990372500db653 Mon Sep 17 00:00:00 2001 From: Penelope Yong Date: Thu, 30 Apr 2026 22:30:00 +0100 Subject: [PATCH 1/3] Add multithreaded assume --- ad.py | 2 +- main.jl | 3 ++- models/multithreaded.jl | 16 ---------------- models/threaded_assume.jl | 14 ++++++++++++++ models/threaded_observe.jl | 13 +++++++++++++ 5 files changed, 30 insertions(+), 18 deletions(-) delete mode 100644 models/multithreaded.jl create mode 100644 models/threaded_assume.jl create mode 100644 models/threaded_observe.jl diff --git a/ad.py b/ad.py index d5a0012..ffc9fbb 100644 --- a/ad.py +++ b/ad.py @@ -78,7 +78,7 @@ def run_ad(args): results = {} - if model_key == "multithreaded": + if model_key.startswith("threaded_"): RUN_JULIA_COMMAND = ["julia", "--threads=4", *JULIA_COMMAND[1:]] else: RUN_JULIA_COMMAND = JULIA_COMMAND diff --git a/main.jl b/main.jl index 3770014..016815e 100644 --- a/main.jl +++ b/main.jl @@ -80,7 +80,8 @@ end # x, y, z, ... are observed variables # although it's hardly a big deal. @include_model "Base Julia features" "control_flow" -@include_model "Base Julia features" "multithreaded" +@include_model "Base Julia features" "threaded_assume" +@include_model "Base Julia features" "threaded_observe" @include_model "Core Turing syntax" "assume_submodel" @include_model "Core Turing syntax" "broadcast_macro" @include_model "Core Turing syntax" "dot_assume" diff --git a/models/multithreaded.jl b/models/multithreaded.jl deleted file mode 100644 index 2f04889..0000000 --- a/models/multithreaded.jl +++ /dev/null @@ -1,16 +0,0 @@ -#= -Most models in ADTests are run with 1 thread. This model is run with 4 threads -to properly demonstrate the compatibility with multithreaded observe -statements. See the docs for more details on multithreading in Turing: -https://turinglang.org/docs/usage/threadsafe-evaluation/ -=# - -@model function multithreaded(x) - a ~ Normal() - Threads.@threads for i in eachindex(x) - x[i] ~ Normal(a) - end -end - -x = randn(100) -model = setthreadsafe(multithreaded(x), true) diff --git a/models/threaded_assume.jl b/models/threaded_assume.jl new file mode 100644 index 0000000..db5bd8a --- /dev/null +++ b/models/threaded_assume.jl @@ -0,0 +1,14 @@ +#= +Note: this example is run with 4 threads +=# + +@model function threaded_assume(x) + a = Vector{Float64}(undef, length(x)) + Threads.@threads for i in eachindex(x) + a[i] ~ Normal() + x[i] ~ Normal(a) + end +end + +x = randn(50) +model = setthreadsafe(threaded_assume(x), true) diff --git a/models/threaded_observe.jl b/models/threaded_observe.jl new file mode 100644 index 0000000..6a21df9 --- /dev/null +++ b/models/threaded_observe.jl @@ -0,0 +1,13 @@ +#= +Note: this model is run with 4 threads +=# +t +@model function threaded_observe(x) + a ~ Normal() + Threads.@threads for i in eachindex(x) + x[i] ~ Normal(a) + end +end + +x = randn(100) +model = setthreadsafe(threaded_observe(x), true) From f6c7191243f548b2fe76d7235282c01186d104f4 Mon Sep 17 00:00:00 2001 From: Penelope Yong Date: Thu, 30 Apr 2026 22:59:12 +0100 Subject: [PATCH 2/3] typo --- models/threaded_observe.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/models/threaded_observe.jl b/models/threaded_observe.jl index 6a21df9..8c565d6 100644 --- a/models/threaded_observe.jl +++ b/models/threaded_observe.jl @@ -1,7 +1,7 @@ #= Note: this model is run with 4 threads =# -t + @model function threaded_observe(x) a ~ Normal() Threads.@threads for i in eachindex(x) From a7f8ecd5a6c5230c8b13af55bd5f8d1e791857fc Mon Sep 17 00:00:00 2001 From: Penelope Yong Date: Thu, 30 Apr 2026 23:08:49 +0100 Subject: [PATCH 3/3] another typo.;; --- models/threaded_assume.jl | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/models/threaded_assume.jl b/models/threaded_assume.jl index db5bd8a..e37f1ba 100644 --- a/models/threaded_assume.jl +++ b/models/threaded_assume.jl @@ -6,7 +6,7 @@ Note: this example is run with 4 threads a = Vector{Float64}(undef, length(x)) Threads.@threads for i in eachindex(x) a[i] ~ Normal() - x[i] ~ Normal(a) + x[i] ~ Normal(a[i]) end end