Files
party2d/validation/summarize_blocked_validation.jl

134 lines
5.8 KiB
Julia

#!/usr/bin/env julia
using CSV
using DataFrames
using Statistics
using Printf
const ECON_VARS = Set(["lrecon_ches", "lrecon_poppa", "lrecon_gps", "lrecon_vparty", "welf_vparty"])
const CULT_VARS = Set(["galtan_ches", "libcon_gps", "immig_vparty", "lgbt_vparty",
"culsup_vparty", "relig_vparty", "gender_vparty"])
function region_for(country::AbstractString)
europe = Set(["AL","AT","BA","BE","BG","BY","CH","CY","CZ","DE","DK","EE","ES","FI","FR","GB","GE","GR","HR","HU","IE","IS","IT","LT","LU","LV","MD","ME","MK","MT","NL","NO","PL","PT","RO","RS","RU","SE","SI","SK","TR","UA"])
latin = Set(["AR","BO","BR","CL","CO","CR","DO","EC","MX","PA","PE","UY"])
country in europe && return "Europe"
country in latin && return "Latin America"
country in Set(["CA", "US"]) && return "North America"
country in Set(["AU", "JP", "KR", "LK", "NZ"]) && return "Asia-Pacific"
return "Other"
end
function wilson(covered::Int, n::Int; z=1.96)
n == 0 && return (missing, missing)
p = covered / n
denom = 1 + z^2 / n
centre = (p + z^2 / (2n)) / denom
half = z * sqrt(p * (1-p) / n + z^2 / (4n^2)) / denom
return (centre - half, centre + half)
end
function safe_cor(x, y)
length(x) < 3 && return missing
std(x) == 0 && return missing
std(y) == 0 && return missing
return cor(x, y)
end
function summarize_group(d::AbstractDataFrame, group_type::String, group::String)
n = nrow(d)
overlapping = count(d.within_latent_interval_95)
lo, hi = wilson(overlapping, n)
DataFrame(
dimension = first(d.dimension), group_type = group_type, group = group,
n = n, parties = length(unique(d.party_id)),
pearson_r = safe_cor(d.expert_value, d.model_value),
mae = mean(d.absolute_error), rmse = sqrt(mean(d.error .^ 2)),
bias_expert_minus_model = mean(d.error),
latent_interval_overlap_95 = overlapping / n,
latent_interval_overlap_ci_lower = lo,
latent_interval_overlap_ci_upper = hi,
mean_interval_width = mean(d.interval_upper .- d.interval_lower),
mean_nearest_text_distance = mean(d.nearest_text_distance)
)
end
function main()
length(ARGS) >= 1 || error("Usage: summarize_blocked_validation.jl MODEL_OUTPUT [RUN_DIR] [OUTPUT_DIR]")
model_file = ARGS[1]
run_dir = length(ARGS) >= 2 ? ARGS[2] : "_local/validation/blocked_party"
output_dir = length(ARGS) >= 3 ? ARGS[3] : "validation/outputs"
mkpath(output_dir)
model = CSV.read(model_file, DataFrame)
expert = CSV.read(joinpath(run_dir, "expert_test.csv"), DataFrame)
text = CSV.read(joinpath(run_dir, "text_data.csv"), DataFrame)
party_col = :party_id in propertynames(model) ? :party_id : :party
model_lookup = Dict{Tuple{Int,Int}, NamedTuple}()
for row in eachrow(model)
model_lookup[(Int(row[party_col]), Int(row.year))] = (
economic_lr = row.economic_lr, economic_lr_q025 = row.economic_lr_q025,
economic_lr_q975 = row.economic_lr_q975, galtan = row.galtan,
galtan_q025 = row.galtan_q025, galtan_q975 = row.galtan_q975)
end
text_years = Dict{Int, Vector{Int}}()
for d in groupby(text, :party)
text_years[Int(first(d.party))] = sort(unique(Int.(d.year)))
end
rows = NamedTuple[]
for row in eachrow(expert)
key = (Int(row.party), Int(row.year))
haskey(model_lookup, key) || continue
dimension = row.var in ECON_VARS ? "economic_lr" : row.var in CULT_VARS ? "galtan" : nothing
isnothing(dimension) && continue
m = model_lookup[key]
if dimension == "economic_lr"
value, lower, upper = m.economic_lr, m.economic_lr_q025, m.economic_lr_q975
else
value, lower, upper = m.galtan, m.galtan_q025, m.galtan_q975
end
years = get(text_years, Int(row.party), Int[])
isempty(years) && continue
distance = minimum(abs.(years .- Int(row.year)))
distance_class = distance == 0 ? "direct text" : distance <= 3 ? "nearby text (1--3 years)" : "distant text (4+ years)"
err = Float64(row.val) - value
push!(rows, (
party_id = Int(row.party), country = String(row.country), region = region_for(String(row.country)),
year = Int(row.year), decade = string(fld(Int(row.year), 10) * 10, "s"),
project = String(row.project), item = String(row.var), dimension = dimension,
expert_value = Float64(row.val), model_value = value,
interval_lower = lower, interval_upper = upper, error = err,
absolute_error = abs(err),
within_latent_interval_95 = lower <= Float64(row.val) <= upper,
nearest_text_distance = distance, text_distance_class = distance_class
))
end
predictions = DataFrame(rows)
isempty(predictions) && error("No held-out expert observations matched the blocked-fit output")
summaries = DataFrame[]
for dimension in sort(unique(predictions.dimension))
d_dim = predictions[predictions.dimension .== dimension, :]
push!(summaries, summarize_group(d_dim, "overall", "All matched held-out ratings"))
for variable in [:decade, :region, :text_distance_class, :project, :item]
for d_group in groupby(d_dim, variable)
push!(summaries, summarize_group(d_group, String(variable), string(first(d_group[!, variable]))))
end
end
end
summary = vcat(summaries...)
predictions_file = joinpath(output_dir, "blocked_validation_predictions.csv")
summary_file = joinpath(output_dir, "blocked_validation_summary.csv")
CSV.write(predictions_file, predictions)
CSV.write(summary_file, summary)
println("Wrote $predictions_file ($(nrow(predictions)) rows)")
println("Wrote $summary_file ($(nrow(summary)) rows)")
println(summary[summary.group_type .== "overall", :])
end
main()