Skip to content
Open
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 docs/Project.toml
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ DataFrames = "a93c6f00-e57d-5684-b7b6-d8193f3e46c0"
Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f"
Documenter = "e30172f5-a6a5-5a46-863b-614d45cd2de4"
FillArrays = "1a297f60-69ca-5386-bcde-b61e274b549b"
Interpolations = "a98d9a8b-a2ab-59e6-89dd-64a1c18fca59"
KernelDensity = "5ab0869b-81aa-558d-bb23-cbf5423bbe9b"
NeuralEstimators = "38f6df31-6b4a-4144-b2af-7ace2da57606"
Plots = "91a5bcdd-55d7-5caf-9e0b-520d859cae80"
Expand All @@ -22,6 +23,7 @@ TuringUtilities = {url = "https://github.com/itsdfish/TuringUtilities.jl"}
Colors = "0.12.0,0.13.0"
DataFrames = "1.0.0"
Documenter = "1"
Interpolations = "0.14.0,0.15.0,0.16.0"
KernelDensity = "0.6.0"
Plots = "1.0.0"
StatsBase = "0.33.0,0.34.0"
Expand Down
3 changes: 3 additions & 0 deletions docs/make.jl
Original file line number Diff line number Diff line change
@@ -1,7 +1,10 @@
using Documenter
using SequentialSamplingModels
using Turing
# load all PlotsExt triggers so the extension (and its docstrings) are available
using Plots
using Interpolations
using KernelDensity

makedocs(
warnonly = true,
Expand Down
78 changes: 47 additions & 31 deletions src/multi_choice_models/DDM.jl
Original file line number Diff line number Diff line change
Expand Up @@ -68,19 +68,30 @@ end
# Wabersich & Vandekerckhove (2014) #
#####################################

function pdf(d::DDM{T}, choice, rt; ϵ::Real = 1.0e-12) where {T <: Real}
@argcheck d.τ < rt
"""
logpdf(d::DDM, choice, rt; ϵ = 1.0e-12)

Log density of `rt` at the boundary given by `choice` (1 = upper, 2 = lower). The density is
computed directly in log space; returns `-Inf` outside the support (`rt ≤ τ`, infinite `rt`,
or an invalid `choice`) rather than throwing, so it can be used safely inside samplers.
"""
function logpdf(d::DDM{T}, choice, rt; ϵ::Real = 1.0e-12) where {T <: Real}
(ν, α, z, τ) = params(d)
LT = typeof(float(zero(promote_type(T, typeof(rt)))))
((choice ≠ 1) && (choice ≠ 2)) && return LT(-Inf)
((rt ≤ τ) || isinf(rt)) && return LT(-Inf)
if choice == 1
(ν, α, z, τ) = params(d)
return _pdf(DDM(-ν, α, 1 - z, τ), rt; ϵ)
return _logpdf_lower(-ν, α, 1 - z, τ, rt; ϵ)
end
return _pdf(d, rt; ϵ)
return _logpdf_lower(ν, α, z, τ, rt; ϵ)
end

# probability density function over the lower boundary
function _pdf(d::DDM, t::Real; ϵ::Real = 1.0e-12)
(ν, α, z, τ) = params(d)
u = (t - τ) / α^2 #use normalized time
pdf(d::DDM, choice, rt; ϵ::Real = 1.0e-12) = exp(logpdf(d, choice, rt; ϵ))

# log probability density function over the lower boundary
function _logpdf_lower(ν, α, z, τ, t; ϵ::Real = 1.0e-12)
Δt = t - τ
u = Δt / α^2 #use normalized time

K_s = 2.0
K_l = 1 / (π * sqrt(u))
Expand All @@ -93,40 +104,45 @@ function _pdf(d::DDM, t::Real; ϵ::Real = 1.0e-12)
K_s = max(2 + sqrt(-2u * log(2ϵ * sqrt(2 * π * u))), sqrt(u) + 1)
end

p = exp((-α * z * ν) - (0.5 * (ν^2) * (t - τ))) / (α^2)
log_p = -α * z * ν - 0.5 * ν^2 * Δt - 2 * log(α)

# decision rule for infinite sum algorithm
if K_s < K_l
return p * _small_time_pdf(u, z, ceil(Int, K_s))
return log_p + _log_small_time_pdf(u, z, ceil(Int, K_s))
end
return p * _large_time_pdf(u, z, ceil(Int, K_l))
log_large = _log_large_time_pdf(u, z, ceil(Int, K_l))
# truncated large-time series can be non-positive; fall back to small-time series
isfinite(log_large) && return log_p + log_large
return log_p + _log_small_time_pdf(u, z, ceil(Int, K_s))
end

# small-time expansion
function _small_time_pdf(u::T, z::T, K::Int) where {T <: Real}
inf_sum = zero(T)

k_series = (-floor(Int, 0.5 * (K - 1))):ceil(Int, 0.5 * (K - 1))
for k ∈ k_series
inf_sum += ((2k + z) * exp(-((2k + z)^2 / (2u))))
# log of the small-time expansion. The k = 0 term has the largest exponent (-z²/2u),
# so it is factored out of the sum to avoid underflow at small u.
function _log_small_time_pdf(u, z, K::Int)
# ((2k + z)² - z²) / 2u = 2k(k + z) / u ≥ 0
term(k) = (2k + z) * exp(-2k * (k + z) / u)
inf_sum = zero(promote_type(typeof(u), typeof(z)))
# sum symmetric ±k pairs from smallest to largest so that they cancel exactly when z = 0.
# This uses at least as many terms as the K terms required by the error bound.
for k ∈ ceil(Int, 0.5 * (K - 1)):-1:1
inf_sum += term(k) + term(-k)
end

return inf_sum / sqrt(2π * u^3)
inf_sum += z
inf_sum ≤ 0 && return oftype(inf_sum, -Inf)
return -z^2 / (2u) + log(inf_sum) - 0.5 * log(2π) - 1.5 * log(u)
end

# large-time expansion
function _large_time_pdf(u::T, z::T, K::Int) where {T <: Real}
inf_sum = zero(T)

# log of the large-time expansion. The k = 1 term has the largest exponent (-π²u/2),
# so it is factored out of the sum to avoid underflow at large u.
function _log_large_time_pdf(u, z, K::Int)
inf_sum = zero(promote_type(typeof(u), typeof(z)))
for k ∈ 1:K
inf_sum += (k * exp(-0.5 * (k^2 * π^2 * u)) * sin(k * π * z))
inf_sum += k * exp(-0.5 * (k^2 - 1) * π^2 * u) * sin(k * π * z)
end

return π * inf_sum
inf_sum ≤ 0 && return oftype(inf_sum, -Inf)
return log(π) - 0.5 * π^2 * u + log(inf_sum)
end

logpdf(d::DDM, choice, rt; ϵ::Real = 1.0e-12) = log(pdf(d, choice, rt; ϵ))

logpdf(d::DDM, data::Tuple) = logpdf(d, data...)

#########################################
Expand All @@ -135,7 +151,7 @@ logpdf(d::DDM, data::Tuple) = logpdf(d, data...)
#########################################

function cdf(d::DDM{T}, choice::Int, rt::Real = 10; ϵ::Real = 1.0e-12) where {T <: Real}
@argcheck d.τ < rt
rt ≤ d.τ && return zero(float(promote_type(T, typeof(rt))))
if choice == 1
(ν, α, z, τ) = params(d)
return _cdf(DDM(-ν, α, 1 - z, τ), rt; ϵ)
Expand Down
1 change: 1 addition & 0 deletions test/Project.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
[deps]
Aqua = "4c88cf16-eb10-579e-8560-4a9242c79595"
Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f"
ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210"
Interpolations = "a98d9a8b-a2ab-59e6-89dd-64a1c18fca59"
JET = "c3a54625-cd67-489e-a8e7-0a5a0ff4e31b"
KernelDensity = "5ab0869b-81aa-558d-bb23-cbf5423bbe9b"
Expand Down
85 changes: 85 additions & 0 deletions test/multi_choice_models/ddm_tests.jl
Original file line number Diff line number Diff line change
Expand Up @@ -292,6 +292,91 @@
@test test_val4 ≈ 0.4450956 atol = 1e-5
end

@safetestset "logpdf" begin
using SequentialSamplingModels
using Test

dist = DDM(; ν = 2.0, α = 1.0, z = 0.5, τ = 0.3)
@test logpdf(dist, 1, 0.5) ≈ log(2.131129) atol = 1e-5
@test logpdf(dist, 2, 0.5) ≈ log(0.2884169) atol = 1e-5
@test pdf(dist, 1, 0.5) ≈ exp(logpdf(dist, 1, 0.5))

# outside of support returns -Inf / 0 rather than throwing
@test logpdf(dist, 1, 0.3) == -Inf
@test logpdf(dist, 2, 0.1) == -Inf
@test logpdf(dist, 1, Inf) == -Inf
@test logpdf(dist, 3, 0.5) == -Inf
@test pdf(dist, 1, 0.1) == 0.0
@test cdf(dist, 1, 0.1) == 0.0

# extreme values stay accurate in log space where the pdf underflows
# reference values computed from the series in 1024-bit BigFloat with 300+ terms
@test logpdf(DDM(2.0, 0.05, 0.5, 0.3), 1, 5.0) ≈ -9279.641942591039
@test pdf(DDM(2.0, 0.05, 0.5, 0.3), 1, 5.0) == 0.0
@test logpdf(DDM(2.0, 3.0, 0.5, 0.3), 1, 0.3001) ≈ -11233.698162868372
@test logpdf(DDM(-8.0, 5.0, 0.99, 0.0), 2, 0.5) ≈ 0.347167902329698
@test logpdf(DDM(8.0, 0.05, 0.01, 0.0), 2, 10.0) ≈ -20055.537212544714
@test logpdf(DDM(0.0, 0.2, 0.3, 0.3), 1, 2.3) ≈ -242.58843967201665

# starting on a boundary gives zero density at that boundary
@test logpdf(DDM(1.0, 0.8, 0.0, 0.3), 2, 0.31) == -Inf
@test logpdf(DDM(1.0, 0.8, 0.0, 0.3), 2, 0.6) == -Inf
@test logpdf(DDM(1.0, 0.8, 1.0, 0.3), 1, 0.31) == -Inf
end

@safetestset "log series" begin
using SequentialSamplingModels: _log_small_time_pdf, _log_large_time_pdf
using Test

# a truncated large-time series can be non-positive, which must map to -Inf
@test _log_large_time_pdf(1e-4, 0.5, 3) == -Inf
# ±k pairs cancel exactly when z = 0
@test _log_small_time_pdf(0.1, 0.0, 5) == -Inf
@test _log_small_time_pdf(0.1, 0.0, 6) == -Inf
# both series agree where they overlap
for z ∈ (0.1, 0.5, 0.9), u ∈ (0.5, 1.0, 2.0)
@test _log_small_time_pdf(u, z, 50) ≈ _log_large_time_pdf(u, z, 50)
end
end

@safetestset "logpdf gradients" begin
using SequentialSamplingModels
using Test
using ForwardDiff

f(θ, c, rt) = logpdf(DDM(θ...), c, rt)
for θ ∈ ([1.0, 0.8, 0.5, 0.3], [0.01, 0.05, 0.1, 0.2], [-3.0, 3.0, 0.9, 0.01])
for c ∈ (1, 2), rt ∈ (0.31, 0.5, 2.0, 10.0)
g = ForwardDiff.gradient(θ -> f(θ, c, rt), θ)
@test all(isfinite, g)
end
end
# gradient is zero (not NaN) outside the support
g = ForwardDiff.gradient(θ -> f(θ, 1, 0.1), [1.0, 0.8, 0.5, 0.3])
@test !any(isnan, g)

# gradients match central finite differences
θ = [0.7, 1.2, 0.4, 0.2]
for c ∈ (1, 2), rt ∈ (0.25, 0.6, 3.0)
g = ForwardDiff.gradient(θ -> f(θ, c, rt), θ)
h = 1e-6
g_fd = map(1:4) do i
e = zeros(4)
e[i] = h
(f(θ .+ e, c, rt) - f(θ .- e, c, rt)) / 2h
end
@test g ≈ g_fd rtol = 1e-5
end

# cdf gradients are finite
F(θ, c, rt) = cdf(DDM(θ...), c, rt)
for θ ∈ ([1.0, 0.8, 0.5, 0.3], [0.01, 0.05, 0.1, 0.2], [-3.0, 3.0, 0.9, 0.01])
for c ∈ (1, 2), rt ∈ (0.31, 0.5, 2.0, 10.0)
@test all(isfinite, ForwardDiff.gradient(θ -> F(θ, c, rt), θ))
end
end
end

@safetestset "simulate" begin
using SequentialSamplingModels
using Test
Expand Down
Loading