From cbc1eaf8fd355ad07e9030a6a4ffbf133c836a6f Mon Sep 17 00:00:00 2001 From: bdeonovic Date: Fri, 21 Jul 2017 10:59:48 -0500 Subject: [PATCH] Add proper ESS/MCSE for discrete and binary and discrete variables --- src/output/chains.jl | 17 +++++++++++++++++ src/output/stats.jl | 33 +++++++++++++++++++++++++++------ 2 files changed, 44 insertions(+), 6 deletions(-) diff --git a/src/output/chains.jl b/src/output/chains.jl index 164c969..8a97b4a 100644 --- a/src/output/chains.jl +++ b/src/output/chains.jl @@ -234,6 +234,23 @@ function indiscretesupport(c::AbstractChains, result end +function inbinarysupport(c::AbstractChains) + nrows, nvars, nchains = size(c.value) + result = Array{Bool}(nvars * (nrows > 0)) + for i in 1:nvars + result[i] = true + result_dict = Set() + for j in 1:nrows, k in 1:nchains + push!(result_dict, c.value[j, i, k]) + if length(result_dict) > 2 + result[i] = false + break + end + end + end + result +end + function link(c::AbstractChains) cc = copy(c.value) for j in 1:length(c.names) diff --git a/src/output/stats.jl b/src/output/stats.jl index af1e527..056a7e3 100644 --- a/src/output/stats.jl +++ b/src/output/stats.jl @@ -83,12 +83,33 @@ function quantile(c::AbstractChains; q::Vector=[0.025, 0.25, 0.5, 0.75, 0.975]) end function summarystats(c::AbstractChains; etype=:bm, args...) - f = x -> [mean(x), std(x), sem(x), mcse(vec(x), etype; args...)] + discrete_flag = indiscretesupport(c) + binary_flag = inbinarysupport(c) + n, p, m = size(c.value) + stats = zeros(Float64, p, 5) + labels = ["Mean", "SD", "Naive SE", "MCSE", "ESS"] - vals = permutedims( - mapslices(x -> f(x), c.value, [1, 3]), - [2, 1, 3] - ) - stats = [vals min.((vals[:, 2] ./ vals[:, 4]).^2, size(c.value, 1))] + + for j in 1:p + if binary_flag[j] + phat = mean(c.value[:,j,:]) + ca = weiss(c.value[:,j,:])[4] + ESS = n / ca + STD = sqrt(phat * (1 - phat)) + SEM = STD / sqrt(n) + MCSE = STD / sqrt(ESS) + stats[j,:] = [phat, STD, SEM, MCSE, ESS] + elseif discrete_flag[j] + x = c.value[:,j,:] + ca = weiss(x)[4] + ESS = n / ca + stats[j,:] = [mean(x), std(x), sem(x), std(x) / sqrt(ESS), ESS] + else + x = c.value[:,j,:] + stats[j,:] = [mean(x), std(x), sem(x), mcse(vec(x), etype; args...), NaN] + stats[j,5] = min((stats[j, 2] ./ stats[j, 4]).^2, n) + end + end + ChainSummary(stats, c.names, labels, header(c)) end